Sfoglia il codice sorgente

feat(qna): /api/v1/qna/advance + /profile 端点与路由注册(LLM 不可用降级 fallback/空画像)

Xiaogang Liao 1 mese fa
parent
commit
dbc324cc6d

+ 2 - 0
cfc-langgraph/src/app.py

@@ -9,6 +9,7 @@ from fastapi.responses import JSONResponse
 
 
 from .schemas.questionnaire import GenerateRequest, GenerateResponse
 from .schemas.questionnaire import GenerateRequest, GenerateResponse
 from .graphs.questionnaire import get_questionnaire_graph
 from .graphs.questionnaire import get_questionnaire_graph
+from .qna.router import router as qna_router
 
 
 router = APIRouter(prefix="/api/v1", tags=["questionnaire"])
 router = APIRouter(prefix="/api/v1", tags=["questionnaire"])
 
 
@@ -42,6 +43,7 @@ async def generate_questionnaire(req: GenerateRequest):
 
 
 app = FastAPI(title="CFC LangGraph 问卷生成服务")
 app = FastAPI(title="CFC LangGraph 问卷生成服务")
 app.include_router(router)
 app.include_router(router)
+app.include_router(qna_router)
 
 
 
 
 @app.get("/health")
 @app.get("/health")

+ 22 - 3
cfc-langgraph/src/qna/graph.py

@@ -40,8 +40,8 @@ def decide_next(llm, scene: dict, history: list, prompt_fn) -> dict:
     prompt = prompt_fn(scene, "", history)  # 知识在 router 中检索后注入,见 retrieve_and_decide
     prompt = prompt_fn(scene, "", history)  # 知识在 router 中检索后注入,见 retrieve_and_decide
     last_err = None
     last_err = None
     for attempt in range(2):
     for attempt in range(2):
-        resp = llm.invoke([{"role": "human", "content": prompt}])
         try:
         try:
+            resp = llm.invoke([{"role": "human", "content": prompt}])
             data = _extract_json(resp.content)
             data = _extract_json(resp.content)
             action = data.get("action")
             action = data.get("action")
             if action == "ask":
             if action == "ask":
@@ -64,8 +64,8 @@ def generate_profile(llm, scene: dict, history: list, kb_context: Optional[list]
     if kb_context:
     if kb_context:
         kb_text = "\n---\n".join(d.get("content", "") for d in kb_context[:5])
         kb_text = "\n---\n".join(d.get("content", "") for d in kb_context[:5])
     prompt = prompt_fn(scene, kb_text, history)
     prompt = prompt_fn(scene, kb_text, history)
-    resp = llm.invoke([{"role": "human", "content": prompt}])
     try:
     try:
+        resp = llm.invoke([{"role": "human", "content": prompt}])
         data = _extract_json(resp.content)
         data = _extract_json(resp.content)
         up = data.get("user_profile", [])
         up = data.get("user_profile", [])
         np = data.get("need_profile", [])
         np = data.get("need_profile", [])
@@ -73,4 +73,23 @@ def generate_profile(llm, scene: dict, history: list, kb_context: Optional[list]
             item["score"] = max(0, min(100, int(item.get("score", 0))))
             item["score"] = max(0, min(100, int(item.get("score", 0))))
         return {"user_profile": up, "need_profile": np}, bool(kb_text)
         return {"user_profile": up, "need_profile": np}, bool(kb_text)
     except Exception:
     except Exception:
-        return {"user_profile": [], "need_profile": []}, False
+        return {"user_profile": [], "need_profile": []}, False
+
+
+def retrieve_and_decide(llm, scene, history, kb_context, prompt_fn) -> dict:
+    """知识注入版出题:把 kb_context 转文本后调用 LLM 决策(供 router 使用)"""
+    kb_text = ""
+    if kb_context:
+        kb_text = "\n---\n".join(d.get("content", "") for d in kb_context[:5])
+    prompt = prompt_fn(scene, kb_text, history)
+    try:
+        resp = llm.invoke([{"role": "human", "content": prompt}])
+        data = _extract_json(resp.content)
+        if data.get("action") == "finish":
+            return {"action": "finish", "reason": data.get("reason", "信息已足够")}
+        q = schemas.Question(**data["question"])
+        return {"action": "ask", "question": q.model_dump(), "reason": data.get("reason", "")}
+    except Exception as e:
+        return {"fallback": True, "reason": f"出题失败: {e}",
+                "question": {"id": f"fb{len(history)+1}", "type": "text",
+                             "text": "请简单描述您最近一周的饮食和作息情况。"}}

+ 72 - 0
cfc-langgraph/src/qna/router.py

@@ -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)