9 Commits
16 changed files with 544 additions and 126 deletions
+14
View File
@@ -4,7 +4,14 @@ from .session import (
ArrowDirection, ArrowDirection,
CardData, CardData,
CardState, CardState,
GuesserMoveEvent,
NextStepRequestedEvent,
PlayerHelloEvent,
RevealSecretEvent,
SessionEvent,
SessionPhase,
SessionStats, SessionStats,
SetterMoveEvent,
StepData, StepData,
StepState, StepState,
) )
@@ -16,7 +23,14 @@ __all__ = [
"ArrowDirection", "ArrowDirection",
"CardData", "CardData",
"CardState", "CardState",
"GuesserMoveEvent",
"NextStepRequestedEvent",
"PlayerHelloEvent",
"RevealSecretEvent",
"SessionEvent",
"SessionPhase",
"SessionStats", "SessionStats",
"SetterMoveEvent",
"StepData", "StepData",
"StepState", "StepState",
] ]
+1 -1
View File
@@ -7,7 +7,7 @@ class PlayerRole(StrEnum):
GUESSER = auto() GUESSER = auto()
@dataclass @dataclass(frozen=True)
class Player: class Player:
name: str name: str
role: PlayerRole role: PlayerRole
+49
View File
@@ -2,6 +2,8 @@ from dataclasses import dataclass
from enum import Enum, StrEnum, auto from enum import Enum, StrEnum, auto
from typing import Optional from typing import Optional
from .common import Player
class ArrowDirection(StrEnum): class ArrowDirection(StrEnum):
LEFT = auto() LEFT = auto()
@@ -16,8 +18,18 @@ class StepState(Enum):
SESSION_FINISHED = auto() SESSION_FINISHED = auto()
class SessionPhase(Enum):
INITIALIZING = auto()
WAITING_FOR_SETTER = auto()
WAITING_FOR_GUESSER = auto()
WAITING_FOR_REVEAL = auto()
CHECKOUT = auto()
FINISHED = auto()
class CardState(Enum): class CardState(Enum):
SETTER_NOT_SET = auto() SETTER_NOT_SET = auto()
SETTER_WAITING = auto()
SETTER_HIDDEN = auto() SETTER_HIDDEN = auto()
SETTER_REVEALED = auto() SETTER_REVEALED = auto()
GUESSER_DISABLED = auto() GUESSER_DISABLED = auto()
@@ -54,3 +66,40 @@ class SessionStats:
accuracy = 0.0 accuracy = 0.0
accuracy = (self.correct / self.total) * 100 accuracy = (self.correct / self.total) * 100
return round(accuracy, 1) return round(accuracy, 1)
@dataclass(frozen=True)
class PlayerHelloEvent:
player: Player
@dataclass(frozen=True)
class SetterMoveEvent:
player: Player
direction: ArrowDirection
@dataclass(frozen=True)
class GuesserMoveEvent:
player: Player
direction: ArrowDirection
@dataclass(frozen=True)
class RevealSecretEvent:
player: Player
direction: ArrowDirection
@dataclass(frozen=True)
class NextStepRequestedEvent:
player: Player
SessionEvent = (
PlayerHelloEvent
| SetterMoveEvent
| GuesserMoveEvent
| RevealSecretEvent
| NextStepRequestedEvent
)
+10 -10
View File
@@ -1,11 +1,11 @@
from .controller import NetworkController # from .controller import NetworkController
from .enums import NetworkMessageKind # from .enums import NetworkMessageKind
from .handlers import SessionNetworkMessageHandler # from .handlers import SessionNetworkMessageHandler
from .message import NetworkMessage # from .message import NetworkMessage
__all__ = [ # __all__ = [
"NetworkController", # "NetworkController",
"NetworkMessageKind", # "NetworkMessageKind",
"SessionNetworkMessageHandler", # "SessionNetworkMessageHandler",
"NetworkMessage", # "NetworkMessage",
] # ]
+71 -41
View File
@@ -1,25 +1,48 @@
from typing import Callable 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
from .enums import NetworkMessageKind from ..state_machines.session.events import SessionEvent
from .message import NetworkMessage from ..state_machines.session.machine import SessionStateMachine
from .message import EventMapper, NetworkMessage, SessionMessageKind
MessageHandler = Callable[[NetworkMessage], None] logger = logging.getLogger(__name__)
class NetworkController: class NetworkController:
def __init__(self, host: str, port: int, username: str, password: str): def __init__(
self._handlers: set[MessageHandler] = set() self,
host: str,
port: int,
username: str,
password: str,
lobby,
session: SessionStateMachine,
):
self._lobby_topic = "mind_reader/lobby"
self._session_topic = "mind_reader/session"
self._lobby = lobby
self._session = session
self._pending_loopbacks: dict[str, NetworkMessage] = {} self._pending_loopbacks: dict[str, NetworkMessage] = {}
self._topic = "mind_reader"
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.on_message = self._on_message
self._client.message_callback_add(
self._lobby_topic, self._on_lobby_message
)
self._client.message_callback_add(
self._session_topic, self._on_session_message
)
self._client.connect(host, port) self._client.connect(host, port)
self._client.subscribe("mind_reader/#")
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()
@@ -28,28 +51,14 @@ class NetworkController:
self._client.loop_stop() self._client.loop_stop()
self._client.disconnect() self._client.disconnect()
def add_handler(self, handler: MessageHandler) -> None:
self._handlers.add(handler)
def remove_handler(self, handler: MessageHandler) -> None:
self._handlers.remove(handler)
def send(self, network_message: NetworkMessage) -> None: def send(self, network_message: NetworkMessage) -> None:
kind = network_message.kind match network_message.kind:
payload = network_message.to_json() case SessionMessageKind():
self._client.publish(
self._session_topic, network_message.to_json()
)
match kind: def send_with_loopback(
case (
NetworkMessageKind.SESSION_PLAYER_MOVED
| NetworkMessageKind.SESSION_SECRET_REVEAL_REQUESTED
| NetworkMessageKind.SESSION_SECRET_REVEALED
| NetworkMessageKind.SESSION_SESSION_FINISHED
| NetworkMessageKind.SESSION_SESSION_STARTED
):
topic = f"{self._topic}/session"
self._client.publish(topic, payload)
def send_with_looback(
self, self,
local_message: NetworkMessage, local_message: NetworkMessage,
network_message: NetworkMessage, network_message: NetworkMessage,
@@ -57,24 +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 _handle(self, network_message: NetworkMessage) -> None: def _pop_loopback(self, message_id: str) -> Optional[NetworkMessage]:
for handler in self._handlers: return self._pending_loopbacks.pop(message_id, None)
handler(network_message)
def _on_message( 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 self, client: mqtt.Client, userdata: None, message: mqtt.MQTTMessage
) -> None: ) -> None:
json_str = message.payload.decode("utf-8") network_message = self._parse_incoming_message(message)
incoming_message = NetworkMessage.from_json(json_str) if network_message is None:
return
event = EventMapper.from_message(network_message)
# if isinstance(event, LobbyEvent):
# self._lobby.process(event)
if incoming_message.id in self._pending_loopbacks: def _on_session_message(
local_message = self._pending_loopbacks.pop(incoming_message.id) self, client: mqtt.Client, userdata: None, message: mqtt.MQTTMessage
self._handle(local_message) ) -> None:
else: network_message = self._parse_incoming_message(message)
self._handle(incoming_message) if network_message is None:
return
event = EventMapper.from_message(network_message)
if isinstance(event, SessionEvent):
self._session.process(event)
-15
View File
@@ -1,15 +0,0 @@
from enum import StrEnum
class NetworkMessageKind(StrEnum):
# ----- Lobby -----
# ---- Session ----
SESSION_SESSION_STARTED = "session_started"
SESSION_PLAYER_MOVED = "player_moved"
SESSION_SECRET_REVEAL_REQUESTED = "secret_reveal_requested"
SESSION_SECRET_REVEALED = "secret_revealed"
SESSION_SESSION_FINISHED = "session_finished"
# ---- System -----
SYSTEM_INC_MSG_PA_ERR = "incoming_message_parsing_error"
-47
View File
@@ -1,47 +0,0 @@
from ..domain import ArrowDirection, Player, PlayerRole
from ..state import SessionState
from .controller import NetworkController
from .enums import NetworkMessageKind
from .message import NetworkMessage
class SessionNetworkMessageHandler:
def __init__(self, network: NetworkController, state: SessionState):
self._network = network
self._state = state
def __call__(self, message: NetworkMessage):
match message.kind:
case NetworkMessageKind.SESSION_SESSION_STARTED:
self._state.start_session()
case NetworkMessageKind.SESSION_PLAYER_MOVED:
player_role = PlayerRole(message.payload["player"]["role"])
direction = ArrowDirection(message.payload["arrow_direction"])
match player_role:
case PlayerRole.SETTER:
self._state.setter_move(direction)
case PlayerRole.GUESSER:
self._state.guesser_move(direction)
case NetworkMessageKind.SESSION_SECRET_REVEAL_REQUESTED:
self._network.send(
NetworkMessage(
NetworkMessageKind.SESSION_SECRET_REVEALED,
{
"player": {
"name": self._state.local_player.name,
"role": self._state.local_player.role,
},
"arrow_direction": self._state.reveal_setter_move(),
},
)
)
case NetworkMessageKind.SESSION_SECRET_REVEALED:
direction = ArrowDirection(message.payload["arrow_direction"])
self._state.step_close(direction)
case NetworkMessageKind.SESSION_SESSION_FINISHED:
pass
+86 -12
View File
@@ -1,34 +1,108 @@
import json import json
import logging
from dataclasses import asdict, dataclass, field from dataclasses import asdict, dataclass, field
from enum import StrEnum, auto
from typing import Any, Optional from typing import Any, Optional
from uuid_extensions import uuid7str from uuid_extensions import uuid7str
from .enums import NetworkMessageKind from ..state_machines.lobby.domain import Player, PlayerRole
from ..state_machines.session.domain import ArrowDirection
from ..state_machines.session.events import (
GuesserMoveEvent,
NextStepRequestedEvent,
PlayerHelloEvent,
RevealSecretEvent,
SessionEvent,
SetterMoveEvent,
)
logger = logging.getLogger(__name__)
class SessionMessageKind(StrEnum):
PLAYER_HELLO = auto()
SETTER_MOVED = auto()
GUESSER_MOVED = auto()
SECRET_REVEALED = auto()
NEXT_STEP_REQUESTED = auto()
NetworkMessageKind = SessionMessageKind
Event = SessionEvent
EVENT_TO_KIND_MAP: dict[type[Event], NetworkMessageKind] = {
PlayerHelloEvent: SessionMessageKind.PLAYER_HELLO,
SetterMoveEvent: SessionMessageKind.SETTER_MOVED,
GuesserMoveEvent: SessionMessageKind.GUESSER_MOVED,
RevealSecretEvent: SessionMessageKind.SECRET_REVEALED,
NextStepRequestedEvent: SessionMessageKind.NEXT_STEP_REQUESTED,
}
KIND_TO_EVENT_MAP = {
kind: event_cls for event_cls, kind in EVENT_TO_KIND_MAP.items()
}
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
payload: dict[str, Any] payload: dict[str, Any] = field(default_factory=dict)
id: str = field(default_factory=uuid7str) id: str = field(default_factory=uuid7str)
def to_json(self) -> str:
data = asdict(self)
data["kind"] = self.kind.value
return json.dumps(data)
@classmethod @classmethod
def from_json(cls, json_str: str) -> "NetworkMessage": def from_json(cls, json_str: str) -> "NetworkMessage":
try: try:
data = json.loads(json_str) data = json.loads(json_str)
return cls( return cls(
id=data["id"],
kind=NetworkMessageKind(data["kind"]), kind=NetworkMessageKind(data["kind"]),
payload=data.get("payload", {}), payload=data.get("payload", {}),
id=data["id"],
) )
except (json.JSONDecodeError, KeyError, ValueError) as e: except (KeyError, ValueError, json.JSONDecodeError) as err:
return cls( logger.error(
kind=NetworkMessageKind.SYSTEM_INC_MSG_PA_ERR, "Failed to deserialize NetworkMessage from JSON: %s", err
payload={"exception": e},
) )
raise
def to_json(self) -> str:
data = {"kind": self.kind.value, "payload": self.payload, "id": self.id}
return json.dumps(data, ensure_ascii=False, default=serializer)
class EventMapper:
@staticmethod
def to_message(event: Event) -> NetworkMessage:
event_type = type(event)
kind = EVENT_TO_KIND_MAP.get(event_type)
if not kind:
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=event.to_payload())
@staticmethod
def from_message(message: NetworkMessage) -> Event | None:
event_cls = KIND_TO_EVENT_MAP.get(message.kind)
if not event_cls:
return None
try:
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)
@@ -0,0 +1,19 @@
from dataclasses import dataclass
from enum import StrEnum, auto
class PlayerRole(StrEnum):
SETTER = auto()
GUESSER = auto()
@dataclass(frozen=True)
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"])
)
@@ -0,0 +1,27 @@
from enum import Enum, StrEnum, auto
class ArrowDirection(StrEnum):
LEFT = auto()
RIGHT = auto()
HIDDEN = auto()
class SessionPhase(Enum):
INITIALIZING = auto()
WAITING_FOR_SETTER = auto()
WAITING_FOR_GUESSER = auto()
WAITING_FOR_REVEAL = auto()
CHECKOUT = auto()
FINISHED = auto()
class CardState(Enum):
SETTER_NOT_SET = auto()
SETTER_WAITING = auto()
SETTER_HIDDEN = auto()
SETTER_REVEALED = auto()
GUESSER_DISABLED = auto()
GUESSER_WAITING = auto()
GUESSER_CORRECT = auto()
GUESSER_WRONG = auto()
@@ -0,0 +1,71 @@
from dataclasses import dataclass
from ..base_event import BaseEvent
from ..lobby.domain import Player
from ..session.domain import ArrowDirection
@dataclass(frozen=True)
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(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(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(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(BaseEvent):
player: Player
@classmethod
def from_payload(cls, p: dict) -> "NextStepRequestedEvent":
return cls(player=Player.from_dict(p))
SessionEvent = (
PlayerHelloEvent
| SetterMoveEvent
| GuesserMoveEvent
| RevealSecretEvent
| NextStepRequestedEvent
)
@@ -0,0 +1,174 @@
from ..lobby.domain import Player, PlayerRole
from ..session.domain import ArrowDirection, CardState, SessionPhase
from ..session.events import (
GuesserMoveEvent,
NextStepRequestedEvent,
PlayerHelloEvent,
RevealSecretEvent,
SessionEvent,
SetterMoveEvent,
)
class SessionStateMachine:
def __init__(
self, setter: Player, guessers: set[Player], sequence_length: int = 20
):
self._setter = setter
self._guessers = guessers
self._pending_players: set[Player] = set(guessers)
self._pending_players.add(setter)
self._sequence_length = sequence_length
self._setter_moves: list[ArrowDirection] = []
self._guessers_moves: dict[Player, list[ArrowDirection]] = {
player: [] for player in guessers
}
self._phase: SessionPhase = SessionPhase.INITIALIZING
def process(self, event: SessionEvent) -> bool:
match (self._phase, event):
case (SessionPhase.INITIALIZING, PlayerHelloEvent(player)):
self._pending_players.remove(player)
if not self._pending_players:
self._phase = SessionPhase.WAITING_FOR_SETTER
self._pending_players = {self._setter}
return True
case (
SessionPhase.WAITING_FOR_SETTER,
SetterMoveEvent(player, direction),
):
self._pending_players.remove(player)
if not self._pending_players:
self._setter_moves.append(direction)
self._phase = SessionPhase.WAITING_FOR_GUESSER
self._pending_players = set(self._guessers)
return True
case (
SessionPhase.WAITING_FOR_GUESSER,
GuesserMoveEvent(player, direction),
):
self._guessers_moves[player].append(direction)
self._pending_players.remove(player)
if not self._pending_players:
self._phase = SessionPhase.WAITING_FOR_REVEAL
self._pending_players = {self._setter}
return True
case (
SessionPhase.WAITING_FOR_REVEAL,
RevealSecretEvent(player, direction),
):
self._pending_players.remove(player)
if not self._pending_players:
self._phase = SessionPhase.CHECKOUT
setter_move_index = len(self._setter_moves) - 1
self._setter_moves[setter_move_index] = direction
self._pending_players = set(self._guessers)
return True
case (SessionPhase.CHECKOUT, NextStepRequestedEvent(player)):
self._pending_players.remove(player)
if not self._pending_players:
if len(self._setter_moves) >= self._sequence_length:
self._phase = SessionPhase.FINISHED
self._pending_players.clear()
else:
self._phase = SessionPhase.WAITING_FOR_SETTER
self._pending_players = {self._setter}
return True
case _:
return False
def get_card_state(self, player: Player, index: int) -> CardState:
match player.role:
case PlayerRole.SETTER:
return self._get_setter_card_state(index)
case PlayerRole.GUESSER:
return self._get_guesser_card_state(player, index)
def _get_setter_card_state(self, index: int) -> CardState:
current_index = self.current_index
if index > current_index:
return CardState.SETTER_NOT_SET
elif index == current_index:
match self._phase:
case SessionPhase.INITIALIZING:
return CardState.SETTER_NOT_SET
case SessionPhase.WAITING_FOR_SETTER:
return CardState.SETTER_WAITING
case (
SessionPhase.WAITING_FOR_GUESSER
| SessionPhase.WAITING_FOR_REVEAL
):
if self._setter_moves[-1] == ArrowDirection.HIDDEN:
return CardState.SETTER_HIDDEN
else:
return CardState.SETTER_REVEALED
case SessionPhase.CHECKOUT | SessionPhase.FINISHED:
return CardState.SETTER_REVEALED
else:
return CardState.SETTER_REVEALED
def _get_guesser_card_state(self, player: Player, index: int) -> CardState:
def _check_direction(
self: SessionStateMachine, player: Player, index: int
) -> CardState:
if self._setter_moves[index] == self._guessers_moves[player][index]:
return CardState.GUESSER_CORRECT
else:
return CardState.GUESSER_WRONG
current_index = self.current_index
if index > current_index:
return CardState.GUESSER_DISABLED
elif index == current_index:
match self._phase:
case (
SessionPhase.INITIALIZING | SessionPhase.WAITING_FOR_SETTER
):
return CardState.GUESSER_DISABLED
case (
SessionPhase.WAITING_FOR_GUESSER
| SessionPhase.WAITING_FOR_REVEAL
):
return CardState.GUESSER_WAITING
case SessionPhase.CHECKOUT | SessionPhase.FINISHED:
return _check_direction(self, player, index)
else:
return _check_direction(self, player, index)
def get_setter_move(self, index: int) -> ArrowDirection | None:
try:
return self._setter_moves[index]
except IndexError:
return None
def get_guesser_move(
self, player: Player, index: int
) -> ArrowDirection | None:
try:
return self._guessers_moves[player][index]
except IndexError:
return None
@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