diff --git a/kalshi/ws/models/communications.py b/kalshi/ws/models/communications.py index b8a270f..d7f85ea 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 be932e8..7f7e013 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 ed5b02a..1a92473 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 cf803e8..1cebdd2 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 95bb62e..5936005 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 d24dc28..7972d57 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 b8358b9..e59e8c0 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)