Skip to content
Open
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
43 changes: 19 additions & 24 deletions massive/websocket/models/models.py
Original file line number Diff line number Diff line change
@@ -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

Expand Down Expand Up @@ -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,
]
41 changes: 41 additions & 0 deletions test_websocket/test_model_types.py
Original file line number Diff line number Diff line change
@@ -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()