diff --git a/lighter/ws_client.py b/lighter/ws_client.py index 14abdf5..0638875 100644 --- a/lighter/ws_client.py +++ b/lighter/ws_client.py @@ -3,6 +3,8 @@ from websockets.client import connect as connect_async from lighter.configuration import Configuration from lighter.endpoint_profiles import join_url +import logging +logger = logging.getLogger(__name__) class WsClient: def __init__( @@ -14,6 +16,8 @@ def __init__( on_order_book_update=print, on_account_update=print, ws_url=None, + on_shutdown=None, + on_unhandled_message=None, ): if ws_url is not None: self.base_url = ws_url.rstrip("/") @@ -35,6 +39,8 @@ def __init__( self.on_order_book_update = on_order_book_update self.on_account_update = on_account_update + self.on_shutdown = on_shutdown + self.on_unhandled_message = on_unhandled_message self.ws = None @@ -57,6 +63,8 @@ def on_message(self, ws, message): elif message_type == "ping": # Respond to ping with pong ws.send(json.dumps({"type": "pong"})) + elif message_type == "shutdown": + self.handle_shutdown(message) else: self.handle_unhandled_message(message) @@ -149,7 +157,17 @@ def handle_update_account(self, message): self.on_account_update(account_id, self.account_states[account_id]) def handle_unhandled_message(self, message): - raise Exception(f"Unhandled message: {message}") + if self.on_unhandled_message: + self.on_unhandled_message(message) + else: + logger.warning("Unhandled message type %r: %s", message.get("type"), message) + + def handle_shutdown(self, message): + close_in_ms = message.get("close_in_ms", 0) + if self.on_shutdown: + self.on_shutdown(close_in_ms) + else: + print(f"Server shutting down, connection closes in {close_in_ms}ms") def on_error(self, ws, error): raise Exception(f"Error: {error}")