|
|
@@ -0,0 +1,72 @@
|
|
|
+"""qna 动态问卷引擎 - HTTP 端点"""
|
|
|
+from fastapi import APIRouter
|
|
|
+from typing import List, Optional
|
|
|
+
|
|
|
+from .schemas import QnaRequest, QnaResponse, ProfileResponse
|
|
|
+from . import graph
|
|
|
+from .graph import retrieve_and_decide
|
|
|
+from .prompts import build_decide_prompt, build_profile_prompt
|
|
|
+from app.rag.retriever import RagRetriever
|
|
|
+
|
|
|
+router = APIRouter(prefix="/api/v1/qna", tags=["qna"])
|
|
|
+
|
|
|
+_retriever: Optional[RagRetriever] = None
|
|
|
+
|
|
|
+
|
|
|
+def _get_retriever() -> Optional[RagRetriever]:
|
|
|
+ global _retriever
|
|
|
+ if _retriever is None:
|
|
|
+ try:
|
|
|
+ _retriever = RagRetriever()
|
|
|
+ except Exception:
|
|
|
+ _retriever = None # 知识库不可用 → 降级
|
|
|
+ return _retriever
|
|
|
+
|
|
|
+
|
|
|
+async def _retrieve(kb_scope: List[str], history: List[dict]) -> list:
|
|
|
+ retriever = _get_retriever()
|
|
|
+ if retriever is None:
|
|
|
+ return []
|
|
|
+ try:
|
|
|
+ query = " ".join([h.get("answer", "") for h in history]) or "肠道健康 饮食习惯"
|
|
|
+ return await retriever.retrieve(query, k=5) # filter 按 scope 由 retriever 支持后接入
|
|
|
+ except Exception:
|
|
|
+ return []
|
|
|
+
|
|
|
+
|
|
|
+def _build_llm():
|
|
|
+ try:
|
|
|
+ from src.llm.client import get_llm
|
|
|
+ return get_llm()
|
|
|
+ except Exception:
|
|
|
+ from langchain_openai import ChatOpenAI
|
|
|
+ import os
|
|
|
+ return ChatOpenAI(model=os.getenv("LLM_MODEL", "deepseek"),
|
|
|
+ api_key=os.getenv("LLM_API_KEY", ""),
|
|
|
+ base_url=os.getenv("LLM_BASE_URL", "https://api.deepseek.com/v1"),
|
|
|
+ temperature=0.7)
|
|
|
+
|
|
|
+
|
|
|
+@router.post("/advance", response_model=QnaResponse)
|
|
|
+async def advance(req: QnaRequest):
|
|
|
+ scene = req.scene.dict()
|
|
|
+ history = [h.dict() for h in req.history]
|
|
|
+ llm = _build_llm()
|
|
|
+ kb_context = await _retrieve(scene.get("kb_scope", []), history)
|
|
|
+ # 知识注入版 decide:检索后带知识出题
|
|
|
+ state = retrieve_and_decide(llm, scene, history, kb_context, build_decide_prompt)
|
|
|
+ if state.get("fallback"):
|
|
|
+ return QnaResponse(action="ask", question=state["question"], reason=state.get("reason"), finished=False)
|
|
|
+ if state["action"] == "finish":
|
|
|
+ return QnaResponse(action="finish", finished=True, reason=state.get("reason", "信息已足够"))
|
|
|
+ return QnaResponse(action="ask", question=state["question"], finished=False)
|
|
|
+
|
|
|
+
|
|
|
+@router.post("/profile", response_model=ProfileResponse)
|
|
|
+async def profile(req: QnaRequest):
|
|
|
+ scene = req.scene.dict()
|
|
|
+ history = [h.dict() for h in req.history]
|
|
|
+ llm = _build_llm()
|
|
|
+ kb_context = await _retrieve(scene.get("kb_scope", []), history)
|
|
|
+ profile_json, kb_used = graph.generate_profile(llm, scene, history, kb_context, build_profile_prompt)
|
|
|
+ return ProfileResponse(profile=profile_json, kb_used=kb_used)
|