diff --git a/lighter/ws_client.py b/lighter/ws_client.py index 14abdf5..5ed2929 100644 --- a/lighter/ws_client.py +++ b/lighter/ws_client.py @@ -42,6 +42,9 @@ def on_message(self, ws, message): if isinstance(message, str): message = json.loads(message) + if not isinstance(message, dict): + return + message_type = message.get("type") if message_type == "connected": @@ -61,7 +64,12 @@ def on_message(self, ws, message): self.handle_unhandled_message(message) async def on_message_async(self, ws, message): - message = json.loads(message) + if isinstance(message, str): + message = json.loads(message) + + if not isinstance(message, dict): + return + message_type = message.get("type") if message_type == "connected": diff --git a/test/test_ws_client.py b/test/test_ws_client.py new file mode 100644 index 0000000..8d49d21 --- /dev/null +++ b/test/test_ws_client.py @@ -0,0 +1,21 @@ +import unittest + +from lighter.ws_client import WsClient + + +class TestWsClient(unittest.IsolatedAsyncioTestCase): + def setUp(self): + self.client = WsClient(order_book_ids=[0]) + + def test_sync_ignores_json_string_message(self): + self.client.on_message(None, '"heartbeat"') + + async def test_async_ignores_json_string_message(self): + await self.client.on_message_async(None, '"heartbeat"') + + async def test_async_ignores_decoded_non_object_message(self): + await self.client.on_message_async(None, ["heartbeat"]) + + +if __name__ == "__main__": + unittest.main()