import asyncio import logging import uuid from fastapi import APIRouter, WebSocket, WebSocketDisconnect from sqlalchemy import select from app.db import async_session_maker from app.llm.orchestrator import run_dm_turn from app.models.game import GameParticipant from app.models.message import Message from app.ws_tickets import consume_ticket logger = logging.getLogger("app.ws_game") router = APIRouter() class ConnectionManager: def __init__(self) -> None: self._rooms: dict[uuid.UUID, set[WebSocket]] = {} self._locks: dict[uuid.UUID, asyncio.Lock] = {} def connect(self, game_id: uuid.UUID, websocket: WebSocket) -> None: self._rooms.setdefault(game_id, set()).add(websocket) def disconnect(self, game_id: uuid.UUID, websocket: WebSocket) -> None: room = self._rooms.get(game_id) if room is not None: room.discard(websocket) if not room: self._rooms.pop(game_id, None) async def broadcast(self, game_id: uuid.UUID, payload: dict) -> None: for ws in list(self._rooms.get(game_id, ())): try: await ws.send_json(payload) except Exception: # noqa: BLE001 — a dead socket shouldn't break the broadcast self.disconnect(game_id, ws) def get_lock(self, game_id: uuid.UUID) -> asyncio.Lock: return self._locks.setdefault(game_id, asyncio.Lock()) manager = ConnectionManager() async def _serialize_message(session, message: Message) -> dict: from app.models.character import Character from app.models.user import User player_name = None character_name = None if message.user_id is not None: user = await session.get(User, message.user_id) player_name = user.name if user else None if message.character_id is not None: character = await session.get(Character, message.character_id) character_name = character.name if character else None return { "id": message.id, "sender_type": message.sender_type, "user_id": str(message.user_id) if message.user_id else None, "player_name": player_name, "character_id": str(message.character_id) if message.character_id else None, "character_name": character_name, "content": message.content, "created_at": message.created_at.isoformat(), } @router.websocket("/ws/games/{game_id}") async def game_websocket(websocket: WebSocket, game_id: uuid.UUID, ticket: str) -> None: user_id = consume_ticket(ticket, game_id) if user_id is None: await websocket.close(code=4401) return await websocket.accept() manager.connect(game_id, websocket) try: while True: data = await websocket.receive_json() if data.get("type") != "message": continue content = (data.get("content") or "").strip() if not content: continue async with manager.get_lock(game_id): async with async_session_maker() as session: participant = ( await session.execute( select(GameParticipant).where( GameParticipant.game_id == game_id, GameParticipant.user_id == user_id, ) ) ).scalar_one_or_none() if participant is None: await websocket.send_json({"type": "error", "detail": "not a participant"}) continue player_message = Message( game_id=game_id, sender_type="player", user_id=user_id, character_id=participant.character_id, content=content, ) session.add(player_message) await session.commit() await session.refresh(player_message) await manager.broadcast( game_id, {"type": "message", "message": await _serialize_message(session, player_message)}, ) await manager.broadcast(game_id, {"type": "typing"}) async def _broadcast_roll(notation: str) -> None: await manager.broadcast(game_id, {"type": "rolling", "notation": notation}) async def _broadcast_game_ended(reason: str) -> None: await manager.broadcast(game_id, {"type": "game_ended", "reason": reason}) try: dm_message = await run_dm_turn( session, game_id, latest_player_message=content, on_roll=_broadcast_roll, on_game_ended=_broadcast_game_ended, ) except Exception: # noqa: BLE001 logger.exception("DM turn failed for game %s", game_id) await manager.broadcast(game_id, {"type": "error", "detail": "dm_turn_failed"}) continue await manager.broadcast( game_id, {"type": "message", "message": await _serialize_message(session, dm_message)}, ) except WebSocketDisconnect: pass finally: manager.disconnect(game_id, websocket)