From 72ebe25e92754c26faeb1fc605be042532ba5d7b Mon Sep 17 00:00:00 2001 From: Pengyi Peng <74917296+pengpengyi92@users.noreply.github.com> Date: Wed, 26 Aug 2026 23:46:40 +0800 Subject: [PATCH] fix: align WebSocketMessage with parsed event type --- massive/websocket/models/models.py | 43 +++++++++++++----------------- test_websocket/test_model_types.py | 41 ++++++++++++++++++++++++++++ 2 files changed, 60 insertions(+), 24 deletions(-) create mode 100644 test_websocket/test_model_types.py diff --git a/massive/websocket/models/models.py b/massive/websocket/models/models.py index cc3d3c16..b4f773a1 100644 --- a/massive/websocket/models/models.py +++ b/massive/websocket/models/models.py @@ -1,4 +1,4 @@ -from typing import Optional, List, Union, NewType +from typing import List, Optional, Union from .common import EventType from ...modelclass import modelclass @@ -444,26 +444,21 @@ def from_dict(d): ) -WebSocketMessage = NewType( - "WebSocketMessage", - List[ - Union[ - EquityAgg, - CurrencyAgg, - EquityTrade, - CryptoTrade, - EquityQuote, - ForexQuote, - CryptoQuote, - Imbalance, - LimitUpLimitDown, - Level2Book, - IndexValue, - LaunchpadValue, - FairMarketValue, - FuturesTrade, - FuturesQuote, - FuturesAgg, - ] - ], -) +WebSocketMessage = Union[ + EquityAgg, + CurrencyAgg, + EquityTrade, + CryptoTrade, + EquityQuote, + ForexQuote, + CryptoQuote, + Imbalance, + LimitUpLimitDown, + Level2Book, + IndexValue, + LaunchpadValue, + FairMarketValue, + FuturesTrade, + FuturesQuote, + FuturesAgg, +] diff --git a/test_websocket/test_model_types.py b/test_websocket/test_model_types.py new file mode 100644 index 00000000..023fbc98 --- /dev/null +++ b/test_websocket/test_model_types.py @@ -0,0 +1,41 @@ +import logging +import unittest + +from massive.websocket import EquityTrade, Market, WebSocketMessage +from massive.websocket.models import parse, parse_single + + +def accept_message(message: WebSocketMessage) -> WebSocketMessage: + """Exercise the public single-message type contract under mypy.""" + return message + + +class WebSocketModelTypesTest(unittest.TestCase): + trade = { + "ev": "T", + "sym": "AAPL", + "x": 10, + "i": "5096", + "z": 3, + "p": 161.87, + "s": 300, + "c": [14, 41], + "t": 1651684192462, + "q": 4009402, + } + + def test_single_parsed_event_matches_message_contract(self): + message = parse_single(self.trade, logging.getLogger(), Market.Stocks) + + self.assertIsInstance(message, EquityTrade) + self.assertIs(accept_message(message), message) + + def test_parse_returns_a_batch_of_messages(self): + messages = parse([self.trade], logging.getLogger(), Market.Stocks) + + self.assertEqual(len(messages), 1) + self.assertIsInstance(messages[0], EquityTrade) + + +if __name__ == "__main__": + unittest.main()