0.5 缩减重构
This commit is contained in:
@@ -1,11 +1,11 @@
|
||||
from fastapi import APIRouter
|
||||
|
||||
from . import chat, health, luckysheet_ws, tasks
|
||||
|
||||
api_router = APIRouter(prefix="/api/v1")
|
||||
api_router.include_router(tasks.router, tags=["tasks"])
|
||||
api_router.include_router(chat.router, tags=["chat"])
|
||||
|
||||
root_router = APIRouter()
|
||||
root_router.include_router(health.router, tags=["health"])
|
||||
root_router.include_router(luckysheet_ws.router)
|
||||
from fastapi import APIRouter
|
||||
|
||||
from . import chat, health, luckysheet_ws, tasks
|
||||
|
||||
api_router = APIRouter(prefix="/api/v1")
|
||||
api_router.include_router(tasks.router, tags=["tasks"])
|
||||
api_router.include_router(chat.router, tags=["chat"])
|
||||
|
||||
root_router = APIRouter()
|
||||
root_router.include_router(health.router, tags=["health"])
|
||||
root_router.include_router(luckysheet_ws.router)
|
||||
|
||||
@@ -1,26 +1,26 @@
|
||||
import asyncio
|
||||
from fastapi import APIRouter
|
||||
from fastapi.responses import StreamingResponse
|
||||
|
||||
from app.deps import AuthDep
|
||||
from app.services.lightrag_service import lightrag_service
|
||||
|
||||
router = APIRouter(prefix="/chat")
|
||||
|
||||
|
||||
@router.get("")
|
||||
async def chat(query: str, document_id: str, auth: AuthDep) -> StreamingResponse: # noqa: ARG001
|
||||
"""
|
||||
SSE 流式占位。后续会调用 LightRAG + OpenAI。
|
||||
当前直接返回 mock 文字,确保前端链路可用。
|
||||
"""
|
||||
|
||||
async def event_stream() -> asyncio.AsyncGenerator[str, None]:
|
||||
reply = await lightrag_service.query(query_text=query, user_id="placeholder-user")
|
||||
chunks = [reply[: len(reply) // 2 or 1], reply[len(reply) // 2 or 1 :]]
|
||||
for chunk in chunks:
|
||||
yield f"data: {chunk}\n\n"
|
||||
await asyncio.sleep(0.1)
|
||||
yield "data: [DONE]\n\n"
|
||||
|
||||
return StreamingResponse(event_stream(), media_type="text/event-stream")
|
||||
import asyncio
|
||||
from fastapi import APIRouter
|
||||
from fastapi.responses import StreamingResponse
|
||||
|
||||
from app.deps import AuthDep
|
||||
from app.services.lightrag_service import lightrag_service
|
||||
|
||||
router = APIRouter(prefix="/chat")
|
||||
|
||||
|
||||
@router.get("")
|
||||
async def chat(query: str, document_id: str, auth: AuthDep) -> StreamingResponse: # noqa: ARG001
|
||||
"""
|
||||
SSE 流式占位。后续会调用 LightRAG + OpenAI。
|
||||
当前直接返回 mock 文字,确保前端链路可用。
|
||||
"""
|
||||
|
||||
async def event_stream() -> asyncio.AsyncGenerator[str, None]:
|
||||
reply = await lightrag_service.query(query_text=query, user_id="placeholder-user")
|
||||
chunks = [reply[: len(reply) // 2 or 1], reply[len(reply) // 2 or 1 :]]
|
||||
for chunk in chunks:
|
||||
yield f"data: {chunk}\n\n"
|
||||
await asyncio.sleep(0.1)
|
||||
yield "data: [DONE]\n\n"
|
||||
|
||||
return StreamingResponse(event_stream(), media_type="text/event-stream")
|
||||
|
||||
@@ -1,195 +1,195 @@
|
||||
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)
|
||||
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)
|
||||
|
||||
@@ -1,29 +1,29 @@
|
||||
from fastapi import APIRouter, HTTPException, status
|
||||
|
||||
from app.deps import AuthDep
|
||||
from app.schemas.tasks import OcrTaskRequest, TaskStatusResponse
|
||||
from app.services.task_tracker import task_tracker
|
||||
from app.workers.tasks import ocr_pipeline
|
||||
|
||||
router = APIRouter(prefix="/tasks")
|
||||
|
||||
|
||||
@router.post("/ocr", response_model=TaskStatusResponse)
|
||||
async def enqueue_ocr_task(payload: OcrTaskRequest, auth: AuthDep) -> TaskStatusResponse:
|
||||
"""记录任务并投递 Celery,阶段 0 返回占位任务。"""
|
||||
task = task_tracker.create_task(user_id=auth.user_id, document_id=payload.document_id, task_type="ocr")
|
||||
ocr_pipeline.delay(
|
||||
task_id=task.task_id,
|
||||
document_id=payload.document_id,
|
||||
file_url=str(payload.file_url),
|
||||
user_id=auth.user_id,
|
||||
)
|
||||
return task
|
||||
|
||||
|
||||
@router.get("/{task_id}", response_model=TaskStatusResponse)
|
||||
async def get_task_status(task_id: str, auth: AuthDep) -> TaskStatusResponse:
|
||||
task = task_tracker.get_task(task_id=task_id, user_id=auth.user_id)
|
||||
if not task:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Task not found")
|
||||
return task
|
||||
from fastapi import APIRouter, HTTPException, status
|
||||
|
||||
from app.deps import AuthDep
|
||||
from app.schemas.tasks import OcrTaskRequest, TaskStatusResponse
|
||||
from app.services.task_tracker import task_tracker
|
||||
from app.workers.tasks import ocr_pipeline
|
||||
|
||||
router = APIRouter(prefix="/tasks")
|
||||
|
||||
|
||||
@router.post("/ocr", response_model=TaskStatusResponse)
|
||||
async def enqueue_ocr_task(payload: OcrTaskRequest, auth: AuthDep) -> TaskStatusResponse:
|
||||
"""记录任务并投递 Celery,阶段 0 返回占位任务。"""
|
||||
task = task_tracker.create_task(user_id=auth.user_id, document_id=payload.document_id, task_type="ocr")
|
||||
ocr_pipeline.delay(
|
||||
task_id=task.task_id,
|
||||
document_id=payload.document_id,
|
||||
file_url=str(payload.file_url),
|
||||
user_id=auth.user_id,
|
||||
)
|
||||
return task
|
||||
|
||||
|
||||
@router.get("/{task_id}", response_model=TaskStatusResponse)
|
||||
async def get_task_status(task_id: str, auth: AuthDep) -> TaskStatusResponse:
|
||||
task = task_tracker.get_task(task_id=task_id, user_id=auth.user_id)
|
||||
if not task:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Task not found")
|
||||
return task
|
||||
|
||||
Reference in New Issue
Block a user