From ef552e99e4ae2c44823352bc9ed161d2a8358928 Mon Sep 17 00:00:00 2001 From: Jeff West Date: Fri, 22 May 2026 08:20:10 -0500 Subject: [PATCH] fix(ws-models): tighten payload ts to AwareDatetime + envelope type Literal `OrderbookDeltaPayload.ts` was missed by the v2.5 #270 WS AwareDatetime sweep that closed the v2.4 #234 REST gap. Widen it to `AwareDatetime | None` so naive RFC3339 strings on the highest-volume WS channel raise `ValidationError` rather than silently passing as `str`, matching `user_orders` and `communications`. Narrow every WS message envelope's `type` field from bare `str` to its channel-specific `Literal[...]` across `OrderbookSnapshotMessage`, `OrderbookDeltaMessage`, `TickerMessage`, `FillMessage`, `UserOrdersMessage`, `MarketLifecycleMessage`, and `CommunicationsMessage`. The default-value form keeps no-arg construction ergonomics but constructing an envelope with a mismatched `type` string now raises `ValidationError`. This also unlocks `pydantic.Discriminator("type")` for the dispatcher's union without a follow-up wire-format change. Five `tests/ws/test_models.py` fixtures sent `"type": "user_orders"` (plural, channel name) where the wire `type` value per the dispatcher table is `"user_order"` (singular). Those fixtures are corrected to match the dispatch table that real frames go through. Closes #331, #353 --- kalshi/ws/models/communications.py | 4 +- kalshi/ws/models/fill.py | 4 +- kalshi/ws/models/market_lifecycle.py | 4 +- kalshi/ws/models/orderbook_delta.py | 8 +- kalshi/ws/models/ticker.py | 4 +- kalshi/ws/models/user_orders.py | 4 +- tests/ws/test_models.py | 132 +++++++++++++++++++++++++-- 7 files changed, 143 insertions(+), 17 deletions(-) diff --git a/kalshi/ws/models/communications.py b/kalshi/ws/models/communications.py index b8a270fc..d7f85ea7 100644 --- a/kalshi/ws/models/communications.py +++ b/kalshi/ws/models/communications.py @@ -1,6 +1,8 @@ """Communications channel message models (RFQ and quote notifications).""" from __future__ import annotations +from typing import Literal + from pydantic import AliasChoices, AwareDatetime, BaseModel, Field from kalshi.types import DollarDecimal, FixedPointCount @@ -152,7 +154,7 @@ class CommunicationsMessage(BaseModel): (RfqCreatedPayload, QuoteCreatedPayload, etc.) to validate msg contents. """ - type: str = "communications" + type: Literal["communications"] = "communications" sid: int seq: int | None = None msg: dict[str, object] diff --git a/kalshi/ws/models/fill.py b/kalshi/ws/models/fill.py index be932e84..7f7e0136 100644 --- a/kalshi/ws/models/fill.py +++ b/kalshi/ws/models/fill.py @@ -1,6 +1,8 @@ """Fill channel message models.""" from __future__ import annotations +from typing import Literal + from pydantic import AliasChoices, BaseModel, Field from kalshi.models.orders import BookSideLiteral, SideLiteral @@ -47,7 +49,7 @@ class FillPayload(BaseModel): class FillMessage(BaseModel): """Fill update message. NO required seq.""" - type: str = "fill" + type: Literal["fill"] = "fill" sid: int seq: int | None = None msg: FillPayload diff --git a/kalshi/ws/models/market_lifecycle.py b/kalshi/ws/models/market_lifecycle.py index ed5b02a1..1a924732 100644 --- a/kalshi/ws/models/market_lifecycle.py +++ b/kalshi/ws/models/market_lifecycle.py @@ -2,7 +2,7 @@ from __future__ import annotations -from typing import Any +from typing import Any, Literal from pydantic import BaseModel @@ -49,7 +49,7 @@ class MarketLifecyclePayload(BaseModel): class MarketLifecycleMessage(BaseModel): """Market lifecycle v2 update message. NO required seq.""" - type: str = "market_lifecycle_v2" + type: Literal["market_lifecycle_v2"] = "market_lifecycle_v2" sid: int seq: int | None = None msg: MarketLifecyclePayload diff --git a/kalshi/ws/models/orderbook_delta.py b/kalshi/ws/models/orderbook_delta.py index cf803e84..1cebdd27 100644 --- a/kalshi/ws/models/orderbook_delta.py +++ b/kalshi/ws/models/orderbook_delta.py @@ -5,7 +5,7 @@ from decimal import Decimal from typing import Annotated, Any, Literal -from pydantic import AliasChoices, BaseModel, BeforeValidator, Field +from pydantic import AliasChoices, AwareDatetime, BaseModel, BeforeValidator, Field from kalshi.types import DollarDecimal, FixedPointCount, _coerce_decimal @@ -107,7 +107,7 @@ class OrderbookDeltaPayload(BaseModel): side: Literal["yes", "no"] client_order_id: str | None = None subaccount: int | None = None - ts: str | None = None + ts: AwareDatetime | None = None # v0.14+ backfill (#162). ts_ms (Unix ms) supersedes ts (RFC3339). ts_ms: int | None = None model_config = {"extra": "allow", "populate_by_name": True} @@ -116,7 +116,7 @@ class OrderbookDeltaPayload(BaseModel): class OrderbookSnapshotMessage(BaseModel): """Full orderbook snapshot, sent on initial subscribe.""" - type: str = "orderbook_snapshot" + type: Literal["orderbook_snapshot"] = "orderbook_snapshot" sid: int seq: int msg: OrderbookSnapshotPayload @@ -126,7 +126,7 @@ class OrderbookSnapshotMessage(BaseModel): class OrderbookDeltaMessage(BaseModel): """Incremental orderbook update.""" - type: str = "orderbook_delta" + type: Literal["orderbook_delta"] = "orderbook_delta" sid: int seq: int msg: OrderbookDeltaPayload diff --git a/kalshi/ws/models/ticker.py b/kalshi/ws/models/ticker.py index 95bb62e9..5936005d 100644 --- a/kalshi/ws/models/ticker.py +++ b/kalshi/ws/models/ticker.py @@ -2,6 +2,8 @@ from __future__ import annotations +from typing import Literal + from pydantic import AliasChoices, BaseModel, Field from kalshi.types import DollarDecimal, FixedPointCount @@ -61,7 +63,7 @@ class TickerPayload(BaseModel): class TickerMessage(BaseModel): """Ticker update message. NO required seq.""" - type: str = "ticker" + type: Literal["ticker"] = "ticker" sid: int seq: int | None = None msg: TickerPayload diff --git a/kalshi/ws/models/user_orders.py b/kalshi/ws/models/user_orders.py index d24dc28f..7972d57c 100644 --- a/kalshi/ws/models/user_orders.py +++ b/kalshi/ws/models/user_orders.py @@ -1,6 +1,8 @@ """User orders channel message models.""" from __future__ import annotations +from typing import Literal + from pydantic import AliasChoices, AwareDatetime, BaseModel, Field from kalshi.models.orders import ( @@ -73,7 +75,7 @@ class UserOrdersPayload(BaseModel): class UserOrdersMessage(BaseModel): """User orders update message. NO required seq.""" - type: str = "user_order" + type: Literal["user_order"] = "user_order" sid: int seq: int | None = None msg: UserOrdersPayload diff --git a/tests/ws/test_models.py b/tests/ws/test_models.py index b8358b92..e59e8c0c 100644 --- a/tests/ws/test_models.py +++ b/tests/ws/test_models.py @@ -152,7 +152,7 @@ def test_delta_with_optional_fields(self) -> None: } msg = OrderbookDeltaMessage.model_validate(raw) assert msg.msg.client_order_id == "my-order" - assert msg.msg.ts == "2026-04-19T18:43:37.662364Z" + assert msg.msg.ts == datetime(2026, 4, 19, 18, 43, 37, 662364, tzinfo=UTC) assert msg.msg.delta == Decimal("-20") # negative delta = removal @@ -394,7 +394,7 @@ def test_market_positions_with_subaccount(self) -> None: class TestUserOrdersModel: def test_parse_user_orders(self) -> None: raw = { - "type": "user_orders", + "type": "user_order", "sid": 5, "msg": user_orders_payload_dict( order_id="ord-001", @@ -415,7 +415,7 @@ def test_parse_user_orders(self) -> None: ), } msg = UserOrdersMessage.model_validate(raw) - assert msg.type == "user_orders" + assert msg.type == "user_order" assert msg.msg.order_id == "ord-001" assert msg.msg.status == "resting" assert msg.msg.is_yes is True @@ -424,7 +424,7 @@ def test_parse_user_orders(self) -> None: def test_user_orders_no_seq(self) -> None: raw = { - "type": "user_orders", + "type": "user_order", "sid": 5, "msg": user_orders_payload_dict(order_id="ord-001"), } @@ -433,7 +433,7 @@ def test_user_orders_no_seq(self) -> None: def test_user_orders_canceled(self) -> None: raw = { - "type": "user_orders", + "type": "user_order", "sid": 5, "msg": user_orders_payload_dict( order_id="ord-002", @@ -1034,7 +1034,7 @@ def test_fill_count_post_position_parse_as_decimal(self) -> None: def test_user_orders_counts_parse_as_decimal(self) -> None: msg = UserOrdersMessage.model_validate( { - "type": "user_orders", + "type": "user_order", "sid": 1, "msg": user_orders_payload_dict( fill_count_fp="3", @@ -1141,7 +1141,7 @@ class TestWsPayloadDatetimeCoercion: def test_user_orders_timestamps_parse_as_datetime(self) -> None: msg = UserOrdersMessage.model_validate( { - "type": "user_orders", + "type": "user_order", "sid": 1, "msg": user_orders_payload_dict( created_time="2026-01-01T00:00:00Z", @@ -1499,3 +1499,121 @@ def test_quote_executed_payload_rejects_naive_executed_ts(self) -> None: "executed_ts": "2026-04-19T18:43:37", } ) + + +class TestIssue331OrderbookDeltaTsAwareDatetime: + """#331: OrderbookDeltaPayload.ts widened to AwareDatetime to match v2.5 #270 sweep.""" + + def test_issue_331_orderbook_delta_ts_is_aware_datetime(self) -> None: + from kalshi.ws.models.orderbook_delta import OrderbookDeltaPayload + + payload = OrderbookDeltaPayload.model_validate( + { + "market_ticker": "T", + "market_id": "x", + "price_dollars": "0.5500", + "delta_fp": "10.00", + "side": "yes", + "ts": "2026-04-19T18:43:37.662364Z", + } + ) + assert isinstance(payload.ts, datetime) + assert payload.ts == datetime(2026, 4, 19, 18, 43, 37, 662364, tzinfo=UTC) + + def test_issue_331_orderbook_delta_ts_rejects_naive(self) -> None: + from kalshi.ws.models.orderbook_delta import OrderbookDeltaPayload + + with pytest.raises(ValidationError, match="timezone"): + OrderbookDeltaPayload.model_validate( + { + "market_ticker": "T", + "market_id": "x", + "price_dollars": "0.5500", + "delta_fp": "10.00", + "side": "yes", + "ts": "2026-04-19T18:43:37", + } + ) + + +class TestIssue353EnvelopeTypeLiteralNarrowing: + """#353: WS message envelope ``type`` narrowed to its channel-specific Literal. + + Construction with a mismatched ``type`` string must raise ValidationError. The + default-value form keeps the no-arg construction ergonomics. + """ + + def test_issue_353_envelope_type_literal_narrowing(self) -> None: + # Each (envelope cls, expected literal value, sample msg) + snapshot_msg = { + "market_ticker": "T", + "market_id": "x", + "yes": [["0.50", "100.00"]], + "no": [["0.45", "150.00"]], + } + delta_msg = { + "market_ticker": "T", + "market_id": "x", + "price_dollars": "0.5500", + "delta_fp": "10.00", + "side": "yes", + } + cases: list[tuple[type, str, dict, dict]] = [ + ( + OrderbookSnapshotMessage, + "orderbook_snapshot", + {"sid": 1, "seq": 1, "msg": snapshot_msg}, + {"sid": 1, "seq": 1, "msg": snapshot_msg, "type": "fill"}, + ), + ( + OrderbookDeltaMessage, + "orderbook_delta", + {"sid": 1, "seq": 1, "msg": delta_msg}, + {"sid": 1, "seq": 1, "msg": delta_msg, "type": "ticker"}, + ), + ( + TickerMessage, + "ticker", + {"sid": 1, "msg": ticker_payload_dict()}, + {"sid": 1, "msg": ticker_payload_dict(), "type": "fill"}, + ), + ( + FillMessage, + "fill", + {"sid": 1, "msg": fill_payload_dict()}, + {"sid": 1, "msg": fill_payload_dict(), "type": "ticker"}, + ), + ( + UserOrdersMessage, + "user_order", + {"sid": 1, "msg": user_orders_payload_dict()}, + {"sid": 1, "msg": user_orders_payload_dict(), "type": "fill"}, + ), + ( + MarketLifecycleMessage, + "market_lifecycle_v2", + { + "sid": 1, + "msg": {"event_type": "created", "market_ticker": "T"}, + }, + { + "sid": 1, + "msg": {"event_type": "created", "market_ticker": "T"}, + "type": "ticker", + }, + ), + ( + CommunicationsMessage, + "communications", + {"sid": 1, "msg": {"id": "rfq-1"}}, + {"sid": 1, "msg": {"id": "rfq-1"}, "type": "ticker"}, + ), + ] + for cls, expected_type, ok_payload, bad_payload in cases: + instance = cls.model_validate(ok_payload) + assert instance.type == expected_type, f"{cls.__name__} default type" + # Field annotation is Literal[], not bare str. + annotation = cls.model_fields["type"].annotation + assert annotation is not str, f"{cls.__name__} still bare str" + with pytest.raises(ValidationError): + cls.model_validate(bad_payload)