Files
mnote/wolai-backend/app/routers/chat.py
T

200 lines
7.6 KiB
Python

from __future__ import annotations
import json
import logging
from typing import AsyncIterator, Optional
from fastapi import APIRouter, HTTPException, status
from fastapi.responses import StreamingResponse
from pydantic import BaseModel
from app.deps import AuthDep
from app.services.lightrag_service import lightrag_service
from app.services.simple_ai_service import simple_ai_service
from app.services.searxng_client import searxng_client
from app.services.supabase_rest import supabase_rest
router = APIRouter(prefix="/chat")
logger = logging.getLogger(__name__)
class ChatRequest(BaseModel):
query: str
document_id: Optional[str] = None
workspace_id: Optional[str] = None
model: Optional[str] = None
use_web_search: bool = False
async def _build_chat_response(
query: str,
auth: AuthDep,
document_id: Optional[str],
workspace_id: Optional[str],
model: Optional[str],
use_web_search: bool,
) -> StreamingResponse:
"""统一封装 GET/POST 的对话逻辑,便于同时支持长文本 POST。"""
if not query or not query.strip():
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="query 不能为空")
resolved_workspace: Optional[str] = None
if workspace_id:
membership = supabase_rest.select_one(
"workspace_members", {"workspace_id": workspace_id, "user_id": auth.user_id}
)
if not membership:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN, detail="无权访问该 workspace"
)
resolved_workspace = workspace_id
if document_id:
document = supabase_rest.select_one(
"documents", {"id": document_id, "user_id": auth.user_id}
)
if not document:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Document not found")
doc_workspace = str(document.get("workspace_id")) if document.get("workspace_id") else None
if resolved_workspace and doc_workspace and resolved_workspace != doc_workspace:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST, detail="workspace 与文档不匹配"
)
resolved_workspace = resolved_workspace or doc_workspace
fallback_result: Optional[dict[str, object]] = None
lightrag_result: Optional[dict[str, object]] = None
web_search_result: Optional[dict[str, object]] = None
fallback_reason: Optional[str] = None
# 联网搜索模式:直接用搜索上下文 + DeepSeek/Ollama 生成
if use_web_search:
search_results = searxng_client.search(query)
context = "\n\n".join(
[f"[{idx+1}] {item['title']}\n{item.get('snippet','')}\n{item['url']}" for idx, item in enumerate(search_results)]
)
web_search_result = await lightrag_service.answer_with_context(
query_text=query,
context=context or "未获取到搜索结果",
prefer_deepseek=True,
)
# 将搜索结果作为引用
if web_search_result is not None:
web_search_result["references"] = search_results
else:
health = await lightrag_service.run_healthcheck()
if not health.get("ok"):
fallback_reason = health.get("error", "LightRAG 未就绪")
else:
try:
lightrag_result = await lightrag_service.stream_answer(
query_text=query,
user_id=auth.user_id,
workspace_id=resolved_workspace,
model_choice=model,
document_id=document_id,
)
except Exception as exc: # pragma: no cover - 运行时保护
fallback_reason = str(exc)
logger.exception("LightRAG stream_answer failed, fallback to simple summary: %s", fallback_reason)
if fallback_reason:
fallback_result = simple_ai_service.summarize_document(
query=query,
user_id=auth.user_id,
document_id=document_id,
workspace_id=resolved_workspace,
)
logger.warning("LightRAG unavailable, use simple summary. reason=%s", fallback_reason)
async def event_stream() -> AsyncIterator[str]:
if web_search_result is not None:
if web_search_result.get("is_streaming") and web_search_result.get("iterator"):
iterator = web_search_result["iterator"]
async for chunk in iterator:
payload = json.dumps({"type": "chunk", "content": chunk})
yield f"data: {payload}\n\n"
else:
payload = json.dumps(
{"type": "chunk", "content": web_search_result.get("content", "")}
)
yield f"data: {payload}\n\n"
references = web_search_result.get("references", []) if isinstance(web_search_result, dict) else []
yield f"data: {json.dumps({'type': 'references', 'data': references})}\n\n"
yield "data: [DONE]\n\n"
return
if fallback_result is not None:
payload = json.dumps(
{
"type": "chunk",
"content": fallback_result.get("content", ""),
}
)
yield f"data: {payload}\n\n"
references = fallback_result.get("references", [])
yield f"data: {json.dumps({'type': 'references', 'data': references})}\n\n"
yield "data: [DONE]\n\n"
return
if lightrag_result is None:
payload = json.dumps({"type": "chunk", "content": "LightRAG 暂不可用"})
yield f"data: {payload}\n\n"
yield f"data: {json.dumps({'type': 'references', 'data': []})}\n\n"
yield "data: [DONE]\n\n"
return
if lightrag_result.get("is_streaming") and lightrag_result.get("iterator"):
iterator = lightrag_result["iterator"]
async for chunk in iterator:
payload = json.dumps({"type": "chunk", "content": chunk})
yield f"data: {payload}\n\n"
else:
payload = json.dumps(
{"type": "chunk", "content": lightrag_result.get("content", "")}
)
yield f"data: {payload}\n\n"
references = lightrag_result.get("references", [])
yield f"data: {json.dumps({'type': 'references', 'data': references})}\n\n"
yield "data: [DONE]\n\n"
return StreamingResponse(event_stream(), media_type="text/event-stream")
@router.get("")
async def chat(
query: str,
auth: AuthDep,
document_id: Optional[str] = None,
workspace_id: Optional[str] = None,
model: Optional[str] = None,
use_web_search: bool = False,
) -> StreamingResponse:
"""
基于 LightRAG 的 SSE 流式回答。
query: 必填问题
document_id: 可选,指定所属 workspace
workspace_id: 可选,直接指定 workspace,优先级高于 document_id
"""
return await _build_chat_response(
query=query,
auth=auth,
document_id=document_id,
workspace_id=workspace_id,
model=model,
use_web_search=use_web_search,
)
@router.post("")
async def chat_post(payload: ChatRequest, auth: AuthDep) -> StreamingResponse:
"""POST 版本,适配长问题与 JSON 传参。"""
return await _build_chat_response(
query=payload.query,
auth=auth,
document_id=payload.document_id,
workspace_id=payload.workspace_id,
model=payload.model,
use_web_search=payload.use_web_search,
)