196 lines
7.0 KiB
Python
196 lines
7.0 KiB
Python
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)
|