from __future__ import annotations import asyncio import json import logging import urllib.parse import zlib from dataclasses import dataclass from typing import Dict, List, Optional import httpx from fastapi import APIRouter, WebSocket, WebSocketDisconnect, status from starlette.concurrency import run_in_threadpool from app.config import settings from app.services.supabase_rest import supabase_rest logger = logging.getLogger(__name__) router = APIRouter(prefix="/ws", tags=["luckysheet"]) @dataclass class LuckysheetClient: websocket: WebSocket user_id: str username: str class LuckysheetConnectionManager: def __init__(self) -> None: self._clients: Dict[str, List[LuckysheetClient]] = {} self._lock = asyncio.Lock() async def add(self, grid_key: str, client: LuckysheetClient) -> None: async with self._lock: self._clients.setdefault(grid_key, []).append(client) logger.debug("Luckysheet client %s joined grid %s", client.user_id, grid_key) async def remove(self, grid_key: str, websocket: WebSocket) -> None: async with self._lock: clients = self._clients.get(grid_key) if not clients: return self._clients[grid_key] = [client for client in clients if client.websocket is not websocket] if not self._clients[grid_key]: self._clients.pop(grid_key, None) logger.debug("Luckysheet connection removed from grid %s", grid_key) async def broadcast_payload( self, grid_key: str, sender: LuckysheetClient, payload: str, event_type: int, ) -> None: clients = self._clients.get(grid_key, []) if not clients: return message = json.dumps( { "data": payload, "id": sender.user_id, "username": sender.username, "type": event_type, }, ensure_ascii=False, ) for client in clients: if client.websocket is sender.websocket: continue try: await client.websocket.send_text(message) except Exception as exc: logger.warning("Failed to forward message to %s: %s", client.user_id, exc) async def broadcast_exit(self, grid_key: str, user_id: str) -> None: clients = self._clients.get(grid_key, []) if not clients: return message = json.dumps({"message": "用户退出", "id": user_id}, ensure_ascii=False) for client in clients: try: await client.websocket.send_text(message) except Exception as exc: logger.warning("Failed to notify client exit: %s", exc) manager = LuckysheetConnectionManager() def _decode_ws_payload(raw_message: str) -> Optional[dict]: try: compressed = raw_message.encode("latin1") inflated = zlib.decompress(compressed) decoded = urllib.parse.unquote(inflated.decode("utf-8")) return json.loads(decoded) except Exception as exc: logger.debug("Failed to decode luckysheet payload: %s", exc) return None async def _fetch_supabase_user(access_token: str) -> Optional[dict]: base_url = settings.supabase_url.rstrip("/") headers = { "apikey": settings.supabase_service_role_key, "Authorization": f"Bearer {access_token}", } try: async with httpx.AsyncClient(timeout=5.0) as client: response = await client.get(f"{base_url}/auth/v1/user", headers=headers) response.raise_for_status() return response.json() except httpx.HTTPError as exc: logger.warning("Supabase auth verification failed: %s", exc) return None async def _fetch_table_by_grid_key(grid_key: str) -> Optional[dict]: try: return await run_in_threadpool(lambda: supabase_rest.select_one("document_tables", {"grid_key": grid_key})) except Exception as exc: logger.error("Failed to query document_tables: %s", exc) return None async def _is_workspace_member(workspace_id: str, user_id: str) -> bool: try: result = await run_in_threadpool( lambda: supabase_rest.select_one("workspace_members", {"workspace_id": workspace_id, "user_id": user_id}), ) return bool(result) except Exception as exc: logger.error("Failed to verify workspace membership: %s", exc) return False @router.websocket("/luckysheet") async def luckysheet_collaboration(websocket: WebSocket) -> None: params = websocket.query_params grid_key = params.get("gridKey") token = params.get("token") raw_user_id = params.get("userid") or params.get("userId") requested_type = params.get("type") or "luckysheet" username = params.get("username") or "" if requested_type != "luckysheet" or not grid_key or not token or not raw_user_id: await websocket.close(code=status.WS_1008_POLICY_VIOLATION, reason="缺少协同参数") return supabase_user = await _fetch_supabase_user(token) if not supabase_user or supabase_user.get("id") != raw_user_id: await websocket.close(code=status.WS_1008_POLICY_VIOLATION, reason="身份验证失败") return table_row = await _fetch_table_by_grid_key(grid_key) if not table_row or not table_row.get("workspace_id"): await websocket.close(code=status.WS_1008_POLICY_VIOLATION, reason="gridKey 无效") return workspace_id = table_row["workspace_id"] if not await _is_workspace_member(workspace_id, raw_user_id): await websocket.close(code=status.WS_1008_POLICY_VIOLATION, reason="无权访问该表格") return fallback_name = ( username or supabase_user.get("user_metadata", {}).get("full_name") or supabase_user.get("email") or f"用户-{raw_user_id[:6]}" ) client = LuckysheetClient(websocket=websocket, user_id=raw_user_id, username=fallback_name) await websocket.accept() await manager.add(grid_key, client) logger.info("Luckysheet client %s connected to %s", raw_user_id, grid_key) try: while True: message = await websocket.receive_text() if message == "rub": continue decoded = _decode_ws_payload(message) if not decoded: continue op_type = decoded.get("t") event_type = 3 if op_type == "mv" else 2 await manager.broadcast_payload(grid_key, client, message, event_type) except WebSocketDisconnect: logger.info("Luckysheet client %s disconnected", raw_user_id) except Exception as exc: logger.error("Luckysheet websocket error: %s", exc) finally: await manager.remove(grid_key, websocket) await manager.broadcast_exit(grid_key, raw_user_id)