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
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)
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)
+28 -34
View File
@@ -1,7 +1,8 @@
import json
import logging
from dataclasses import asdict, dataclass, field
from enum import StrEnum, auto
from typing import Any
from typing import Any, Optional
from uuid_extensions import uuid7str
@@ -16,6 +17,8 @@ from ..state_machines.session.events import (
SetterMoveEvent,
)
logger = logging.getLogger(__name__)
class SessionMessageKind(StrEnum):
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)
class NetworkMessage:
kind: NetworkMessageKind
@@ -50,17 +61,22 @@ class NetworkMessage:
@classmethod
def from_json(cls, json_str: str) -> "NetworkMessage":
try:
data = json.loads(json_str)
return cls(
kind=NetworkMessageKind(data["kind"]),
payload=data.get("payload", {}),
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:
data = asdict(self)
data["kind"] = self.kind.value
return json.dumps(data, ensure_ascii=False)
data = {"kind": self.kind.value, "payload": self.payload, "id": self.id}
return json.dumps(data, ensure_ascii=False, default=serializer)
class EventMapper:
@@ -70,9 +86,14 @@ class EventMapper:
kind = EVENT_TO_KIND_MAP.get(event_type)
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
def from_message(message: NetworkMessage) -> Event | None:
@@ -82,33 +103,6 @@ class EventMapper:
return None
try:
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)
return event_cls.from_payload(message.payload)
except (KeyError, ValueError, TypeError):
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:
name: str
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 ..base_event import BaseEvent
from ..lobby.domain import Player
from ..session.domain import ArrowDirection
@dataclass(frozen=True)
class PlayerHelloEvent:
class PlayerHelloEvent(BaseEvent):
player: Player
@classmethod
def from_payload(cls, p: dict) -> "PlayerHelloEvent":
return cls(player=Player.from_dict(p))
@dataclass(frozen=True)
class SetterMoveEvent:
class SetterMoveEvent(BaseEvent):
player: Player
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)
class GuesserMoveEvent:
class GuesserMoveEvent(BaseEvent):
player: Player
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)
class RevealSecretEvent:
class RevealSecretEvent(BaseEvent):
player: Player
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)
class NextStepRequestedEvent:
class NextStepRequestedEvent(BaseEvent):
player: Player
@classmethod
def from_payload(cls, p: dict) -> "NextStepRequestedEvent":
return cls(player=Player.from_dict(p))
SessionEvent = (
PlayerHelloEvent
@@ -93,7 +93,7 @@ class SessionStateMachine:
return self._get_guesser_card_state(player, index)
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:
return CardState.SETTER_NOT_SET
@@ -127,7 +127,7 @@ class SessionStateMachine:
else:
return CardState.GUESSER_WRONG
current_index = self.get_guesser_step_number(player) - 1
current_index = self.current_index
if index > current_index:
return CardState.GUESSER_DISABLED
@@ -149,15 +149,6 @@ class SessionStateMachine:
else:
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:
try:
return self._setter_moves[index]
@@ -172,10 +163,12 @@ class SessionStateMachine:
except IndexError:
return None
def get_last_setter_move(self) -> ArrowDirection | None:
index = self.get_setter_step_number() - 1
return self.get_setter_move(index)
def get_last_guesser_move(self, player: Player) -> ArrowDirection | None:
index = self.get_guesser_step_number(player) - 1
return self.get_guesser_move(player, index)
@property
def current_index(self) -> int:
setter_step_number = len(self._setter_moves)
if setter_step_number:
if self._phase == SessionPhase.WAITING_FOR_SETTER:
return setter_step_number
else:
return setter_step_number - 1
return 0