diff --git a/src/mind_reader/network/controller.py b/src/mind_reader/network/controller.py index e2e33dc..fbbc716 100644 --- a/src/mind_reader/network/controller.py +++ b/src/mind_reader/network/controller.py @@ -1,3 +1,6 @@ +import logging +from typing import Optional + import paho.mqtt.client as mqtt from paho.mqtt.enums import CallbackAPIVersion @@ -5,6 +8,8 @@ from ..state_machines.session.events import SessionEvent from ..state_machines.session.machine import SessionStateMachine from .message import EventMapper, NetworkMessage, SessionMessageKind +logger = logging.getLogger(__name__) + class NetworkController: def __init__( @@ -26,10 +31,6 @@ class NetworkController: self._client = mqtt.Client(CallbackAPIVersion.VERSION2) self._client.username_pw_set(username, password) - self._client.connect(host, port) - - self._client.subscribe(self._lobby_topic) - self._client.subscribe(self._session_topic) self._client.message_callback_add( self._lobby_topic, self._on_lobby_message @@ -38,6 +39,11 @@ class NetworkController: self._session_topic, self._on_session_message ) + self._client.connect(host, port) + + self._client.subscribe(self._lobby_topic) + self._client.subscribe(self._session_topic) + def start(self) -> None: self._client.loop_start() @@ -47,12 +53,12 @@ class NetworkController: def send(self, network_message: NetworkMessage) -> None: match network_message.kind: - case SessionMessageKind: + case SessionMessageKind(): self._client.publish( self._session_topic, network_message.to_json() ) - def send_with_looback( + def send_with_loopback( self, local_message: NetworkMessage, network_message: NetworkMessage, @@ -60,39 +66,45 @@ class NetworkController: if local_message.id != network_message.id: raise ValueError( f"Mismatched message IDs in send_with_loopback: " - f"local_message.id='{local_message.id}' vs network_message.id='{network_message.id}'" + f"local_message.id='{local_message.id}'" + f"vs network_message.id='{network_message.id}'" ) self._pending_loopbacks[local_message.id] = local_message self.send(network_message) - def _check_loopbacks(self, id: str) -> NetworkMessage | None: - if id in self._pending_loopbacks: - return self._pending_loopbacks.pop(id) - return None + def _pop_loopback(self, message_id: str) -> Optional[NetworkMessage]: + return self._pending_loopbacks.pop(message_id, None) + + def _parse_incoming_message( + self, message: mqtt.MQTTMessage + ) -> Optional[NetworkMessage]: + try: + json_str = message.payload.decode() + network_message = NetworkMessage.from_json(json_str) + except Exception as err: + logger.error("Failed to parse incoming MQTT message: %s", err) + return None + + loopback_message = self._pop_loopback(network_message.id) + return loopback_message or network_message def _on_lobby_message( self, client: mqtt.Client, userdata: None, message: mqtt.MQTTMessage ) -> None: - json_str = message.payload.decode() - network_message = NetworkMessage.from_json(json_str) - - loopback_message = self._check_loopbacks(network_message.id) - if loopback_message: - network_message = loopback_message - + network_message = self._parse_incoming_message(message) + if network_message is None: + return event = EventMapper.from_message(network_message) + # if isinstance(event, LobbyEvent): + # self._lobby.process(event) def _on_session_message( self, client: mqtt.Client, userdata: None, message: mqtt.MQTTMessage ) -> None: - json_str = message.payload.decode() - network_message = NetworkMessage.from_json(json_str) - - loopback_message = self._check_loopbacks(network_message.id) - if loopback_message: - network_message = loopback_message - + network_message = self._parse_incoming_message(message) + if network_message is None: + return event = EventMapper.from_message(network_message) if isinstance(event, SessionEvent): self._session.process(event)