"""
WebSocket 連線管理器(階段 3c)。

設計重點:
  - 單 process in-memory pub/sub(階段 3 用 SQLite 沒 Redis,夠 5-15 人團隊)
  - 同一個使用者可以開多個 WS(手機 + 電腦同時登入)
  - HTTP API 改完資料後 → 呼叫 manager.broadcast() → 所有連線收到事件
  - 客戶端自己 filter 事件相關性(channel_id 對得上才更新 UI)

未來 Stage 5 上 NAS 多 process 時,把這個換成 Redis Pub/Sub 即可,
HTTP API 那層的呼叫介面不用改。
"""

from collections import defaultdict
from typing import Any
from uuid import UUID

from fastapi import WebSocket


class ConnectionManager:
    """管理所有開著的 WebSocket 連線,提供廣播功能。"""

    def __init__(self) -> None:
        # user_id → [WebSocket, WebSocket, ...](同一人多裝置)
        self._connections: dict[UUID, list[WebSocket]] = defaultdict(list)

    async def connect(self, websocket: WebSocket, user_id: UUID) -> None:
        """接受連線並登記,呼叫前 caller 自己處理 accept()。"""
        self._connections[user_id].append(websocket)

    def disconnect(self, websocket: WebSocket, user_id: UUID) -> None:
        """從登記表移除一條連線(連線意外掉的時候被呼叫)。"""
        if user_id not in self._connections:
            return
        try:
            self._connections[user_id].remove(websocket)
        except ValueError:
            pass
        if not self._connections[user_id]:
            del self._connections[user_id]

    async def broadcast(
        self,
        event: dict[str, Any],
        exclude_user_id: UUID | None = None,
    ) -> None:
        """把事件送給所有(或除了 exclude_user_id 以外的)連線。

        各 endpoint 改資料後呼叫,例如:
            await manager.broadcast({"type": "message_new", "channel_id": "...", "message": {...}})

        傳送失敗的連線會被自動清掉(可能對方已關瀏覽器但伺服器還沒收到 close)。
        """
        dead: list[tuple[UUID, WebSocket]] = []
        for user_id, ws_list in list(self._connections.items()):
            if exclude_user_id is not None and user_id == exclude_user_id:
                continue
            for ws in list(ws_list):
                try:
                    await ws.send_json(event)
                except Exception:
                    dead.append((user_id, ws))
        for user_id, ws in dead:
            self.disconnect(ws, user_id)

    def connection_count(self) -> int:
        """目前在線連線數(debug / monitoring 用)。"""
        return sum(len(v) for v in self._connections.values())

    def online_user_ids(self) -> set[UUID]:
        """目前在線使用者 ID 集合。"""
        return set(self._connections.keys())


# 全域單例(整個 app 共用)
manager = ConnectionManager()
