Compare commits
4
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
6017c2d757 | ||
|
|
53be71316c | ||
|
|
473b03c1bb | ||
|
|
a07d94bcc4 |
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user