2025-12-06 16:47:17 +08:00
|
|
|
from __future__ import annotations
|
|
|
|
|
|
|
|
|
|
import json
|
2025-12-08 19:56:24 +08:00
|
|
|
import logging
|
2025-12-06 16:47:17 +08:00
|
|
|
from typing import AsyncIterator, Optional
|
|
|
|
|
|
|
|
|
|
from fastapi import APIRouter, HTTPException, status
|
2025-11-23 10:55:04 +08:00
|
|
|
from fastapi.responses import StreamingResponse
|
2025-12-08 19:56:24 +08:00
|
|
|
from pydantic import BaseModel
|
2025-11-23 10:55:04 +08:00
|
|
|
|
|
|
|
|
from app.deps import AuthDep
|
|
|
|
|
from app.services.lightrag_service import lightrag_service
|
2025-12-08 19:56:24 +08:00
|
|
|
from app.services.simple_ai_service import simple_ai_service
|
|
|
|
|
from app.services.searxng_client import searxng_client
|
2025-12-06 16:47:17 +08:00
|
|
|
from app.services.supabase_rest import supabase_rest
|
2025-11-23 10:55:04 +08:00
|
|
|
|
|
|
|
|
router = APIRouter(prefix="/chat")
|
2025-12-08 19:56:24 +08:00
|
|
|
logger = logging.getLogger(__name__)
|
2025-11-23 10:55:04 +08:00
|
|
|
|
|
|
|
|
|
2025-12-08 19:56:24 +08:00
|
|
|
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(
|
2025-12-06 16:47:17 +08:00
|
|
|
query: str,
|
|
|
|
|
auth: AuthDep,
|
2025-12-08 19:56:24 +08:00
|
|
|
document_id: Optional[str],
|
|
|
|
|
workspace_id: Optional[str],
|
|
|
|
|
model: Optional[str],
|
|
|
|
|
use_web_search: bool,
|
2025-12-06 16:47:17 +08:00
|
|
|
) -> StreamingResponse:
|
2025-12-08 19:56:24 +08:00
|
|
|
"""统一封装 GET/POST 的对话逻辑,便于同时支持长文本 POST。"""
|
2025-12-06 16:47:17 +08:00
|
|
|
if not query or not query.strip():
|
|
|
|
|
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="query 不能为空")
|
2025-11-23 10:55:04 +08:00
|
|
|
|
2025-12-08 19:56:24 +08:00
|
|
|
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
|
|
|
|
|
|
2025-12-06 16:47:17 +08:00
|
|
|
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")
|
2025-12-08 19:56:24 +08:00
|
|
|
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
|
2025-12-06 16:47:17 +08:00
|
|
|
|
2025-12-08 19:56:24 +08:00
|
|
|
fallback_result: Optional[dict[str, object]] = None
|
|
|
|
|
lightrag_result: Optional[dict[str, object]] = None
|
|
|
|
|
web_search_result: Optional[dict[str, object]] = None
|
2025-12-06 16:47:17 +08:00
|
|
|
|
2025-12-08 19:56:24 +08:00
|
|
|
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)
|
2025-12-06 16:47:17 +08:00
|
|
|
|
|
|
|
|
async def event_stream() -> AsyncIterator[str]:
|
2025-12-08 19:56:24 +08:00
|
|
|
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
|
|
|
|
|
|
2025-12-06 16:47:17 +08:00
|
|
|
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"
|
2025-11-23 10:55:04 +08:00
|
|
|
yield "data: [DONE]\n\n"
|
|
|
|
|
|
|
|
|
|
return StreamingResponse(event_stream(), media_type="text/event-stream")
|
2025-12-08 19:56:24 +08:00
|
|
|
|
|
|
|
|
|
|
|
|
|
@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,
|
|
|
|
|
)
|