8 Commits
6 changed files with 145 additions and 88 deletions
+36 -24
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)
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 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)
+28 -34
View File
@@ -1,7 +1,8 @@
import json import json
import logging
from dataclasses import asdict, dataclass, field from dataclasses import asdict, dataclass, field
from enum import StrEnum, auto from enum import StrEnum, auto
from typing import Any from typing import Any, Optional
from uuid_extensions import uuid7str from uuid_extensions import uuid7str
@@ -16,6 +17,8 @@ from ..state_machines.session.events import (
SetterMoveEvent, SetterMoveEvent,
) )
logger = logging.getLogger(__name__)
class SessionMessageKind(StrEnum): class SessionMessageKind(StrEnum):
PLAYER_HELLO = auto() PLAYER_HELLO = auto()
@@ -42,6 +45,14 @@ KIND_TO_EVENT_MAP = {
} }
def serializer(obj: Any) -> Any:
if isinstance(obj, StrEnum):
return obj.value
raise TypeError(
f"Object of type {type(obj).__name__} is not JSON serializable"
)
@dataclass(frozen=True) @dataclass(frozen=True)
class NetworkMessage: class NetworkMessage:
kind: NetworkMessageKind kind: NetworkMessageKind
@@ -50,17 +61,22 @@ class NetworkMessage:
@classmethod @classmethod
def from_json(cls, json_str: str) -> "NetworkMessage": def from_json(cls, json_str: str) -> "NetworkMessage":
try:
data = json.loads(json_str) data = json.loads(json_str)
return cls( return cls(
kind=NetworkMessageKind(data["kind"]), kind=NetworkMessageKind(data["kind"]),
payload=data.get("payload", {}), payload=data.get("payload", {}),
id=data["id"], id=data["id"],
) )
except (KeyError, ValueError, json.JSONDecodeError) as err:
logger.error(
"Failed to deserialize NetworkMessage from JSON: %s", err
)
raise
def to_json(self) -> str: def to_json(self) -> str:
data = asdict(self) data = {"kind": self.kind.value, "payload": self.payload, "id": self.id}
data["kind"] = self.kind.value return json.dumps(data, ensure_ascii=False, default=serializer)
return json.dumps(data, ensure_ascii=False)
class EventMapper: class EventMapper:
@@ -70,9 +86,14 @@ class EventMapper:
kind = EVENT_TO_KIND_MAP.get(event_type) kind = EVENT_TO_KIND_MAP.get(event_type)
if not kind: if not kind:
raise ValueError(f"Event {event_type} not registered") logger.error(
"Attempted to serialize unregistered event: %s", event_type
)
raise ValueError(
f"Event {event_type} is not registered in EVENT_TO_KIND_MAP"
)
return NetworkMessage(kind=kind, payload=asdict(event)) return NetworkMessage(kind=kind, payload=event.to_payload())
@staticmethod @staticmethod
def from_message(message: NetworkMessage) -> Event | None: def from_message(message: NetworkMessage) -> Event | None:
@@ -82,33 +103,6 @@ class EventMapper:
return None return None
try: try:
payload = message.payload return event_cls.from_payload(message.payload)
player = Player(
name=payload["player"]["name"],
role=PlayerRole(payload["player"]["role"]),
)
match message.kind:
case NetworkMessageKind.PLAYER_HELLO:
return PlayerHelloEvent(player=player)
case NetworkMessageKind.SETTER_MOVED:
return SetterMoveEvent(
player=player,
direction=ArrowDirection(payload["direction"]),
)
case NetworkMessageKind.GUESSER_MOVED:
return GuesserMoveEvent(
player=player,
direction=ArrowDirection(payload["direction"]),
)
case NetworkMessageKind.SECRET_REVEALED:
return RevealSecretEvent(
player=player,
direction=ArrowDirection(payload["direction"]),
)
case NetworkMessageKind.NEXT_STEP_REQUESTED:
return NextStepRequestedEvent(player=player)
except (KeyError, ValueError, TypeError): except (KeyError, ValueError, TypeError):
return None return None
@@ -0,0 +1,22 @@
from dataclasses import asdict, dataclass, is_dataclass
from enum import StrEnum
from typing import Any
def _factory(d: list[tuple[str, Any]]) -> dict[str, Any]:
def _convert(obj: Any) -> Any:
if isinstance(obj, StrEnum):
return obj.value
if isinstance(obj, list | tuple):
return [_convert(item) for item in obj]
if isinstance(obj, dict):
return {k: _convert(v) for k, v in obj.items()}
return obj
return {k: _convert(v) for k, v in d}
@dataclass(frozen=True)
class BaseEvent:
def to_payload(self) -> dict[str, Any]:
return asdict(self, dict_factory=_factory)
@@ -11,3 +11,9 @@ class PlayerRole(StrEnum):
class Player: class Player:
name: str name: str
role: PlayerRole role: PlayerRole
@classmethod
def from_dict(cls, d: dict) -> "Player":
return cls(
name=d["player"]["name"], role=PlayerRole(d["player"]["role"])
)
@@ -1,36 +1,66 @@
from dataclasses import dataclass from dataclasses import dataclass
from ..base_event import BaseEvent
from ..lobby.domain import Player from ..lobby.domain import Player
from ..session.domain import ArrowDirection from ..session.domain import ArrowDirection
@dataclass(frozen=True) @dataclass(frozen=True)
class PlayerHelloEvent: class PlayerHelloEvent(BaseEvent):
player: Player player: Player
@classmethod
def from_payload(cls, p: dict) -> "PlayerHelloEvent":
return cls(player=Player.from_dict(p))
@dataclass(frozen=True) @dataclass(frozen=True)
class SetterMoveEvent: class SetterMoveEvent(BaseEvent):
player: Player player: Player
direction: ArrowDirection direction: ArrowDirection
@classmethod
def from_payload(cls, p: dict) -> "SetterMoveEvent":
return cls(
player=Player.from_dict(p),
direction=ArrowDirection(p["arrow_direction"]),
)
@dataclass(frozen=True) @dataclass(frozen=True)
class GuesserMoveEvent: class GuesserMoveEvent(BaseEvent):
player: Player player: Player
direction: ArrowDirection direction: ArrowDirection
@classmethod
def from_payload(cls, p: dict) -> "GuesserMoveEvent":
return cls(
player=Player.from_dict(p),
direction=ArrowDirection(p["arrow_direction"]),
)
@dataclass(frozen=True) @dataclass(frozen=True)
class RevealSecretEvent: class RevealSecretEvent(BaseEvent):
player: Player player: Player
direction: ArrowDirection direction: ArrowDirection
@classmethod
def from_payload(cls, p: dict) -> "RevealSecretEvent":
return cls(
player=Player.from_dict(p),
direction=ArrowDirection(p["arrow_direction"]),
)
@dataclass(frozen=True) @dataclass(frozen=True)
class NextStepRequestedEvent: class NextStepRequestedEvent(BaseEvent):
player: Player player: Player
@classmethod
def from_payload(cls, p: dict) -> "NextStepRequestedEvent":
return cls(player=Player.from_dict(p))
SessionEvent = ( SessionEvent = (
PlayerHelloEvent PlayerHelloEvent
@@ -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