Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 3 additions & 1 deletion kalshi/ws/models/communications.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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]
Expand Down
4 changes: 3 additions & 1 deletion kalshi/ws/models/fill.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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
Expand Down
4 changes: 2 additions & 2 deletions kalshi/ws/models/market_lifecycle.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@

from __future__ import annotations

from typing import Any
from typing import Any, Literal

from pydantic import BaseModel

Expand Down Expand Up @@ -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
Expand Down
8 changes: 4 additions & 4 deletions kalshi/ws/models/orderbook_delta.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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}
Expand All @@ -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
Expand All @@ -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
Expand Down
4 changes: 3 additions & 1 deletion kalshi/ws/models/ticker.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,8 @@

from __future__ import annotations

from typing import Literal

from pydantic import AliasChoices, BaseModel, Field

from kalshi.types import DollarDecimal, FixedPointCount
Expand Down Expand Up @@ -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
Expand Down
4 changes: 3 additions & 1 deletion kalshi/ws/models/user_orders.py
Original file line number Diff line number Diff line change
@@ -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 (
Expand Down Expand Up @@ -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
Expand Down
132 changes: 125 additions & 7 deletions tests/ws/test_models.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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",
Expand All @@ -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
Expand All @@ -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"),
}
Expand All @@ -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",
Expand Down Expand Up @@ -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",
Expand Down Expand Up @@ -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",
Expand Down Expand Up @@ -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[<expected>], 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)
Loading