feat: почти рабочий вариант session

This commit is contained in:
2026-08-18 03:53:34 +03:00
parent a7fcf3dda8
commit b78081d999
16 changed files with 189 additions and 227 deletions
+6 -4
View File
@@ -1,9 +1,11 @@
from .controller import NetworkController
from .protocol import NetworkMesssage
from .session_network_controller import SessionNetworkController
from .enums import NetworkMessageKind
from .handlers import SessionNetworkMessageHandler
from .message import NetworkMessage
__all__ = [
"SessionNetworkController",
"NetworkController",
"NetworkMesssage",
"NetworkMessageKind",
"SessionNetworkMessageHandler",
"NetworkMessage",
]
+46 -138
View File
@@ -1,150 +1,58 @@
import socket
import struct
import threading
from typing import Callable, Optional
from typing import Callable
import wx
import paho.mqtt.client as mqtt
from paho.mqtt.enums import CallbackAPIVersion
from .protocol import NetworkMesssage
from .enums import NetworkMessageKind
from .message import NetworkMessage
OnMessageCallback = Callable[[NetworkMesssage], None]
OnNoticeCallback = Callable[[str], None]
MessageHandler = Callable[[NetworkMessage], None]
def receive_exact(sock: socket.socket, n: int) -> Optional[bytes]:
data = bytearray()
class NetworkController:
def __init__(self, host: str, port: int, username: str, password: str):
self._handlers: set[MessageHandler] = set()
while len(data) < n:
packet = sock.recv(n - len(data))
if not packet:
return None
data.extend(packet)
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.connect(host, port)
self._client.subscribe("mind_reader/#")
return bytes(data)
class NetworkController(threading.Thread):
def __init__(
self,
on_message_callback: OnMessageCallback,
on_notice_callback: Optional[OnNoticeCallback] = None,
):
super().__init__(daemon=True)
self.on_message_callback = on_message_callback
self.on_notice_callback = on_notice_callback
self.sock: Optional[socket.socket] = None
self.conn: Optional[socket.socket] = None
self.is_running: bool = False
self.is_server: bool = False
def serve(self, host: str = "0.0.0.0", port: int = 8994) -> None:
self.is_server = True
self.sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
self.sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
self.sock.bind((host, port))
self.sock.listen(1)
self.is_running = True
self.start()
def connect(self, host: str, port: int = 8994) -> bool:
self.is_server = False
try:
self.conn = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
self.conn.connect((host, port))
self.is_running = True
self._notify("Connected to host")
self.start()
return True
except Exception as e:
self._notify(f"Connection failed: {e}")
return False
def send(self, event: str, payload: dict = {}) -> None:
if not self.conn or not self.is_running:
return
m = NetworkMesssage(event, payload)
try:
self.conn.sendall(m.encode())
except Exception as e:
self._notify(f"Send error: {e}")
self._close_active_connection()
def run(self) -> None:
HEADER_SIZE = 4
while self.is_running:
if self.is_server and self.sock and not self.conn:
self._notify("Waiting for peer to connect...")
try:
self.conn, addr = self.sock.accept()
self._notify(f"Peer connected from {addr[0]}:{addr[1]}")
except Exception:
break
while self.is_running and self.conn:
try:
header_bytes = receive_exact(self.conn, HEADER_SIZE)
if not header_bytes:
self._notify("Peer disconnected")
break
body_size = struct.unpack(">I", header_bytes)[0]
body_bytes = receive_exact(self.conn, body_size)
if not body_bytes:
self._notify("Peer disconnected unexpectedly")
break
json_str = body_bytes.decode("utf-8")
m = NetworkMesssage.from_json(json_str)
wx.CallAfter(self.on_message_callback, m)
except Exception as e:
self._notify(f"Read error: {e}")
break
self._close_active_connection()
if not self.is_server:
break
self.stop()
def _close_active_connection(self) -> None:
if self.conn:
try:
self.conn.shutdown(socket.SHUT_RDWR)
self.conn.close()
except Exception:
pass
finally:
self.conn = None
def _notify(self, notice: str) -> None:
if self.on_notice_callback:
wx.CallAfter(self.on_notice_callback, notice)
def start(self) -> None:
self._client.loop_start()
def stop(self) -> None:
self.is_running = False
self._client.loop_stop()
self._client.disconnect()
self._close_active_connection()
def add_handler(self, handler: MessageHandler) -> None:
self._handlers.add(handler)
if self.sock:
try:
self.sock.close()
except Exception:
pass
finally:
self.sock = None
def remove_handler(self, handler: MessageHandler) -> None:
self._handlers.remove(handler)
self._notify("Connection closed")
def send(self, network_message: NetworkMessage) -> None:
kind = network_message.kind
payload = network_message.to_json()
match kind:
case (
NetworkMessageKind.SESSION_PLAYER_MOVED
| NetworkMessageKind.SESSION_SESSION_FINISHED
| NetworkMessageKind.SESSION_SESSION_STARTED
):
topic = f"{self._topic}/session"
self._client.publish(topic, payload)
def _handle(self, network_message: NetworkMessage) -> None:
for handler in self._handlers:
handler(network_message)
def _on_message(
self, client: mqtt.Client, userdata: None, message: mqtt.MQTTMessage
) -> None:
json_str = message.payload.decode("utf-8")
network_message = NetworkMessage.from_json(json_str)
self._handle(network_message)
+13
View File
@@ -0,0 +1,13 @@
from enum import StrEnum
class NetworkMessageKind(StrEnum):
# ----- Lobby -----
# ---- Session ----
SESSION_SESSION_STARTED = "session_started"
SESSION_PLAYER_MOVED = "player_moved"
SESSION_SESSION_FINISHED = "session_finished"
# ---- System -----
SYSTEM_INC_MSG_PA_ERR = "incoming_message_parsing_error"
+27
View File
@@ -0,0 +1,27 @@
from ..domain import ArrowDirection, Player, PlayerRole
from ..state import SessionState
from .enums import NetworkMessageKind
from .message import NetworkMessage
class SessionNetworkMessageHandler:
def __init__(self, state: SessionState):
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_SESSION_FINISHED:
pass
+34
View File
@@ -0,0 +1,34 @@
import json
from dataclasses import asdict, dataclass, field
from typing import Any, Optional
from uuid_extensions import uuid7str
from .enums import NetworkMessageKind
@dataclass(frozen=True)
class NetworkMessage:
kind: NetworkMessageKind
payload: dict[str, Any]
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", {}),
)
except (json.JSONDecodeError, KeyError, ValueError) as e:
return cls(
kind=NetworkMessageKind.SYSTEM_INC_MSG_PA_ERR,
payload={"exception": e},
)
-46
View File
@@ -1,46 +0,0 @@
import itertools
import json
import struct
from dataclasses import dataclass, field
_id_generator = itertools.count(0)
@dataclass
class NetworkMesssage:
event: str
payload: dict[str, str] = field(default_factory=dict)
id: int = field(default_factory=lambda: next(_id_generator))
version: str = "1.0"
def to_json(self) -> str:
return json.dumps(
{
"version": self.version,
"id": self.id,
"event": self.event,
"payload": self.payload,
},
ensure_ascii=False,
)
@classmethod
def from_json(cls, json_str: str) -> "NetworkMesssage":
data = json.loads(json_str)
return cls(
event=data.get("event", "unknown"),
payload=data.get("payload", {}),
id=data.get("id", -1),
version=data.get("version", "unknown"),
)
def encode(self) -> bytes:
raw_bytes = self.to_json().encode("utf-8")
length_prefix = struct.pack(">I", len(raw_bytes))
return length_prefix + raw_bytes
if __name__ == "__main__":
for i in range(10):
message = NetworkMesssage("none")
print(message.encode())
@@ -1,2 +0,0 @@
class SessionNetworkController:
pass