Compare commits

4 Commits
2 changed files with 48 additions and 43 deletions
+37 -25
View File
@@ -1,3 +1,6 @@
import logging
from typing import Optional
import paho.mqtt.client as mqtt import paho.mqtt.client as mqtt
from paho.mqtt.enums import CallbackAPIVersion 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 ..state_machines.session.machine import SessionStateMachine
from .message import EventMapper, NetworkMessage, SessionMessageKind from .message import EventMapper, NetworkMessage, SessionMessageKind
logger = logging.getLogger(__name__)
class NetworkController: class NetworkController:
def __init__( def __init__(
@@ -26,10 +31,6 @@ class NetworkController:
self._client = mqtt.Client(CallbackAPIVersion.VERSION2) self._client = mqtt.Client(CallbackAPIVersion.VERSION2)
self._client.username_pw_set(username, password) 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._client.message_callback_add(
self._lobby_topic, self._on_lobby_message self._lobby_topic, self._on_lobby_message
@@ -38,6 +39,11 @@ class NetworkController:
self._session_topic, self._on_session_message 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: def start(self) -> None:
self._client.loop_start() self._client.loop_start()
@@ -47,12 +53,12 @@ class NetworkController:
def send(self, network_message: NetworkMessage) -> None: def send(self, network_message: NetworkMessage) -> None:
match network_message.kind: match network_message.kind:
case SessionMessageKind: case SessionMessageKind():
self._client.publish( self._client.publish(
self._session_topic, network_message.to_json() self._session_topic, network_message.to_json()
) )
def send_with_looback( def send_with_loopback(
self, self,
local_message: NetworkMessage, local_message: NetworkMessage,
network_message: NetworkMessage, network_message: NetworkMessage,
@@ -60,39 +66,45 @@ class NetworkController:
if local_message.id != network_message.id: if local_message.id != network_message.id:
raise ValueError( raise ValueError(
f"Mismatched message IDs in send_with_loopback: " 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._pending_loopbacks[local_message.id] = local_message
self.send(network_message) self.send(network_message)
def _check_loopbacks(self, id: str) -> NetworkMessage | None: def _pop_loopback(self, message_id: str) -> Optional[NetworkMessage]:
if id in self._pending_loopbacks: return self._pending_loopbacks.pop(message_id, None)
return self._pending_loopbacks.pop(id)
return 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( def _on_lobby_message(
self, client: mqtt.Client, userdata: None, message: mqtt.MQTTMessage self, client: mqtt.Client, userdata: None, message: mqtt.MQTTMessage
) -> None: ) -> None:
json_str = message.payload.decode() network_message = self._parse_incoming_message(message)
network_message = NetworkMessage.from_json(json_str) if network_message is None:
return
loopback_message = self._check_loopbacks(network_message.id)
if loopback_message:
network_message = loopback_message
event = EventMapper.from_message(network_message) event = EventMapper.from_message(network_message)
# if isinstance(event, LobbyEvent):
# self._lobby.process(event)
def _on_session_message( def _on_session_message(
self, client: mqtt.Client, userdata: None, message: mqtt.MQTTMessage self, client: mqtt.Client, userdata: None, message: mqtt.MQTTMessage
) -> None: ) -> None:
json_str = message.payload.decode() network_message = self._parse_incoming_message(message)
network_message = NetworkMessage.from_json(json_str) if network_message is None:
return
loopback_message = self._check_loopbacks(network_message.id)
if loopback_message:
network_message = loopback_message
event = EventMapper.from_message(network_message) event = EventMapper.from_message(network_message)
if isinstance(event, SessionEvent): if isinstance(event, SessionEvent):
self._session.process(event) self._session.process(event)
@@ -93,7 +93,7 @@ class SessionStateMachine:
return self._get_guesser_card_state(player, index) return self._get_guesser_card_state(player, index)
def _get_setter_card_state(self, index: int) -> CardState: def _get_setter_card_state(self, index: int) -> CardState:
current_index = self.get_setter_step_number() - 1 current_index = self.current_index
if index > current_index: if index > current_index:
return CardState.SETTER_NOT_SET return CardState.SETTER_NOT_SET
@@ -127,7 +127,7 @@ class SessionStateMachine:
else: else:
return CardState.GUESSER_WRONG return CardState.GUESSER_WRONG
current_index = self.get_guesser_step_number(player) - 1 current_index = self.current_index
if index > current_index: if index > current_index:
return CardState.GUESSER_DISABLED return CardState.GUESSER_DISABLED
@@ -149,15 +149,6 @@ class SessionStateMachine:
else: else:
return _check_direction(self, player, index) return _check_direction(self, player, index)
def _get_step_number(self, enum) -> int:
return len(enum)
def get_setter_step_number(self) -> int:
return self._get_step_number(self._setter_moves)
def get_guesser_step_number(self, player: Player) -> int:
return self._get_step_number(self._guessers_moves[player])
def get_setter_move(self, index: int) -> ArrowDirection | None: def get_setter_move(self, index: int) -> ArrowDirection | None:
try: try:
return self._setter_moves[index] return self._setter_moves[index]
@@ -172,10 +163,12 @@ class SessionStateMachine:
except IndexError: except IndexError:
return None return None
def get_last_setter_move(self) -> ArrowDirection | None: @property
index = self.get_setter_step_number() - 1 def current_index(self) -> int:
return self.get_setter_move(index) setter_step_number = len(self._setter_moves)
if setter_step_number:
def get_last_guesser_move(self, player: Player) -> ArrowDirection | None: if self._phase == SessionPhase.WAITING_FOR_SETTER:
index = self.get_guesser_step_number(player) - 1 return setter_step_number
return self.get_guesser_move(player, index) else:
return setter_step_number - 1
return 0