Compare commits
9
Commits
5c7966755f
...
main
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4533789432 | ||
|
|
b1132715e1 | ||
|
|
878b83dc1c | ||
|
|
28030412d0 | ||
|
|
6017c2d757 | ||
|
|
53be71316c | ||
|
|
473b03c1bb | ||
|
|
a07d94bcc4 | ||
|
|
9a4c98fb76 |
@@ -4,7 +4,14 @@ from .session import (
|
||||
ArrowDirection,
|
||||
CardData,
|
||||
CardState,
|
||||
GuesserMoveEvent,
|
||||
NextStepRequestedEvent,
|
||||
PlayerHelloEvent,
|
||||
RevealSecretEvent,
|
||||
SessionEvent,
|
||||
SessionPhase,
|
||||
SessionStats,
|
||||
SetterMoveEvent,
|
||||
StepData,
|
||||
StepState,
|
||||
)
|
||||
@@ -16,7 +23,14 @@ __all__ = [
|
||||
"ArrowDirection",
|
||||
"CardData",
|
||||
"CardState",
|
||||
"GuesserMoveEvent",
|
||||
"NextStepRequestedEvent",
|
||||
"PlayerHelloEvent",
|
||||
"RevealSecretEvent",
|
||||
"SessionEvent",
|
||||
"SessionPhase",
|
||||
"SessionStats",
|
||||
"SetterMoveEvent",
|
||||
"StepData",
|
||||
"StepState",
|
||||
]
|
||||
|
||||
@@ -7,7 +7,7 @@ class PlayerRole(StrEnum):
|
||||
GUESSER = auto()
|
||||
|
||||
|
||||
@dataclass
|
||||
@dataclass(frozen=True)
|
||||
class Player:
|
||||
name: str
|
||||
role: PlayerRole
|
||||
|
||||
@@ -2,6 +2,8 @@ from dataclasses import dataclass
|
||||
from enum import Enum, StrEnum, auto
|
||||
from typing import Optional
|
||||
|
||||
from .common import Player
|
||||
|
||||
|
||||
class ArrowDirection(StrEnum):
|
||||
LEFT = auto()
|
||||
@@ -16,8 +18,18 @@ class StepState(Enum):
|
||||
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):
|
||||
SETTER_NOT_SET = auto()
|
||||
SETTER_WAITING = auto()
|
||||
SETTER_HIDDEN = auto()
|
||||
SETTER_REVEALED = auto()
|
||||
GUESSER_DISABLED = auto()
|
||||
@@ -54,3 +66,40 @@ class SessionStats:
|
||||
accuracy = 0.0
|
||||
accuracy = (self.correct / self.total) * 100
|
||||
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
|
||||
)
|
||||
|
||||
@@ -1,11 +1,11 @@
|
||||
from .controller import NetworkController
|
||||
from .enums import NetworkMessageKind
|
||||
from .handlers import SessionNetworkMessageHandler
|
||||
from .message import NetworkMessage
|
||||
# from .controller import NetworkController
|
||||
# from .enums import NetworkMessageKind
|
||||
# from .handlers import SessionNetworkMessageHandler
|
||||
# from .message import NetworkMessage
|
||||
|
||||
__all__ = [
|
||||
"NetworkController",
|
||||
"NetworkMessageKind",
|
||||
"SessionNetworkMessageHandler",
|
||||
"NetworkMessage",
|
||||
]
|
||||
# __all__ = [
|
||||
# "NetworkController",
|
||||
# "NetworkMessageKind",
|
||||
# "SessionNetworkMessageHandler",
|
||||
# "NetworkMessage",
|
||||
# ]
|
||||
|
||||
@@ -1,25 +1,48 @@
|
||||
from typing import Callable
|
||||
import logging
|
||||
from typing import Optional
|
||||
|
||||
import paho.mqtt.client as mqtt
|
||||
from paho.mqtt.enums import CallbackAPIVersion
|
||||
|
||||
from .enums import NetworkMessageKind
|
||||
from .message import NetworkMessage
|
||||
from ..state_machines.session.events import SessionEvent
|
||||
from ..state_machines.session.machine import SessionStateMachine
|
||||
from .message import EventMapper, NetworkMessage, SessionMessageKind
|
||||
|
||||
MessageHandler = Callable[[NetworkMessage], None]
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class NetworkController:
|
||||
def __init__(self, host: str, port: int, username: str, password: str):
|
||||
self._handlers: set[MessageHandler] = set()
|
||||
def __init__(
|
||||
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._topic = "mind_reader"
|
||||
self._client = mqtt.Client(CallbackAPIVersion.VERSION2)
|
||||
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.subscribe("mind_reader/#")
|
||||
|
||||
self._client.subscribe(self._lobby_topic)
|
||||
self._client.subscribe(self._session_topic)
|
||||
|
||||
def start(self) -> None:
|
||||
self._client.loop_start()
|
||||
@@ -28,28 +51,14 @@ class NetworkController:
|
||||
self._client.loop_stop()
|
||||
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:
|
||||
kind = network_message.kind
|
||||
payload = network_message.to_json()
|
||||
match network_message.kind:
|
||||
case SessionMessageKind():
|
||||
self._client.publish(
|
||||
self._session_topic, network_message.to_json()
|
||||
)
|
||||
|
||||
match kind:
|
||||
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(
|
||||
def send_with_loopback(
|
||||
self,
|
||||
local_message: NetworkMessage,
|
||||
network_message: NetworkMessage,
|
||||
@@ -57,24 +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 _handle(self, network_message: NetworkMessage) -> None:
|
||||
for handler in self._handlers:
|
||||
handler(network_message)
|
||||
def _pop_loopback(self, message_id: str) -> Optional[NetworkMessage]:
|
||||
return self._pending_loopbacks.pop(message_id, None)
|
||||
|
||||
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
|
||||
) -> None:
|
||||
json_str = message.payload.decode("utf-8")
|
||||
incoming_message = NetworkMessage.from_json(json_str)
|
||||
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)
|
||||
|
||||
if incoming_message.id in self._pending_loopbacks:
|
||||
local_message = self._pending_loopbacks.pop(incoming_message.id)
|
||||
self._handle(local_message)
|
||||
else:
|
||||
self._handle(incoming_message)
|
||||
def _on_session_message(
|
||||
self, client: mqtt.Client, userdata: None, message: mqtt.MQTTMessage
|
||||
) -> None:
|
||||
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)
|
||||
|
||||
@@ -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"
|
||||
@@ -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
|
||||
@@ -1,34 +1,108 @@
|
||||
import json
|
||||
import logging
|
||||
from dataclasses import asdict, dataclass, field
|
||||
from enum import StrEnum, auto
|
||||
from typing import Any, Optional
|
||||
|
||||
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)
|
||||
class NetworkMessage:
|
||||
kind: NetworkMessageKind
|
||||
payload: dict[str, Any]
|
||||
payload: dict[str, Any] = field(default_factory=dict)
|
||||
id: str = field(default_factory=uuid7str)
|
||||
|
||||
def to_json(self) -> str:
|
||||
data = asdict(self)
|
||||
data["kind"] = self.kind.value
|
||||
return json.dumps(data)
|
||||
|
||||
@classmethod
|
||||
def from_json(cls, json_str: str) -> "NetworkMessage":
|
||||
try:
|
||||
data = json.loads(json_str)
|
||||
return cls(
|
||||
id=data["id"],
|
||||
kind=NetworkMessageKind(data["kind"]),
|
||||
payload=data.get("payload", {}),
|
||||
id=data["id"],
|
||||
)
|
||||
except (json.JSONDecodeError, KeyError, ValueError) as e:
|
||||
return cls(
|
||||
kind=NetworkMessageKind.SYSTEM_INC_MSG_PA_ERR,
|
||||
payload={"exception": e},
|
||||
except (KeyError, ValueError, json.JSONDecodeError) as err:
|
||||
logger.error(
|
||||
"Failed to deserialize NetworkMessage from JSON: %s", err
|
||||
)
|
||||
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
|
||||
Reference in New Issue
Block a user