| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322 |
- from typing import TypedDict, Optional, Literal
- from langgraph.graph import StateGraph, START, END
- from langgraph.checkpoint.memory import MemorySaver
- from langchain_openai import ChatOpenAI, OpenAIEmbeddings
- from langchain_core.messages import SystemMessage, HumanMessage
- from app.tools.java_client import JavaClient
- from app.tools.product_tools import (
- search_product_by_keyword,
- search_article_by_keyword,
- )
- from app.memory.store import MemoryManager
- from app.config import settings
- import logging
- import json
- logger = logging.getLogger(__name__)
- class HealthButlerState(TypedDict):
- query: str
- user_id: int
- conversation_id: str
- family_id: int | None
- child_id: int | None
- report_id: int | None
- focus: str | None # gut / nutrition / overall / chronic
- kb_context: dict | None
- knowledge_results: list[dict]
- answer: str | None
- tasks: list[dict]
- sources: list[dict]
- messages: list | None
- HEALTH_BUTLER_PROMPT = """你是一个家庭AI健康管家, 为用户提供专业、温暖的个性化健康建议。
- 你的能力:
- 1. 基于健康知识库(菌属功能、指标解读、营养素、食物)提供科学建议
- 2. 解读健康报告数据, 结合知识库给出改善方案
- 3. 根据菌群检测结果推荐精准营养方案
- 4. 提供可执行的健康微行动
- 回答原则:
- 1. 用中文, 像家庭医生一样亲切但专业
- 2. 引用知识库中的菌属/指标/食物信息作为证据
- 3. 优先使用知识库数据, 其次使用通用健康知识
- 4. 需要产品推荐时使用搜索工具
- 5. 生成 [TASK: {"title": "任务名", "dimension": "身", "points": 10, "frequency": "每天"}] 标记创建健康行动任务
- 6. 超出知识库范围的慢性病/急症问题, 建议咨询专业医生
- 7. 不要编造医学建议, 不夸大效果
- """
- def create_health_butler_graph():
- """创建 AI 健康管家 StateGraph — 知识检索 + LLM 建议 + 记忆"""
- llm = ChatOpenAI(
- model=settings.llm_model,
- api_key=settings.llm_api_key,
- base_url=settings.llm_base_url,
- temperature=settings.llm_temperature,
- )
- embeddings = OpenAIEmbeddings(
- model=settings.embedding_model,
- api_key=settings.effective_embedding_api_key,
- base_url=settings.effective_embedding_base_url,
- )
- memory_mgr = MemoryManager(embeddings)
- java = JavaClient()
- llm_with_tools = llm.bind_tools([
- search_product_by_keyword,
- search_article_by_keyword,
- ])
- builder = StateGraph(HealthButlerState)
- async def load_knowledge(state: HealthButlerState) -> dict:
- """阶段1: 从 Java 后端加载健康知识库上下文"""
- context = {}
- try:
- client = await java._get_client()
- resp = await client.post(
- "/api/health/knowledge/list",
- json={},
- )
- data = resp.json()
- if data.get("code") == 200:
- items = data.get("data", [])
- # 按 itemType 分组
- by_type = {}
- for item in items:
- t = item.get("itemType", "unknown")
- if t not in by_type:
- by_type[t] = []
- by_type[t].append(item)
- context["kb_items"] = len(items)
- context["kb_by_type"] = {k: len(v) for k, v in by_type.items()}
- context["kb_loaded"] = True
- else:
- context["kb_loaded"] = False
- except Exception as e:
- logger.warning("加载健康知识库失败: %s", e)
- context["kb_loaded"] = False
- return {"kb_context": context}
- async def search_knowledge(state: HealthButlerState) -> dict:
- """阶段2: 根据 query 意图精准搜索知识库"""
- query = state["query"]
- results = []
- # 提取 query 中可能的菌属名/指标名/食物名关键词
- # 策略: 对 query 做简单的关键词匹配, 调用 Java 知识库逐个查询
- keywords = _extract_health_keywords(query)
- if keywords:
- client = await java._get_client()
- for kw in keywords[:5]: # 最多查询5个关键词
- # 优先查 bacteria 类型
- for item_type in ["bacteria", "indicator", "food", "nutrient"]:
- try:
- resp = await client.post(
- "/api/health/knowledge/query",
- json={"itemType": item_type, "itemName": kw},
- )
- data = resp.json()
- if data.get("code") == 200 and data.get("data"):
- results.append({
- "type": item_type,
- "keyword": kw,
- "data": data["data"],
- })
- break # 命中即跳出
- except Exception:
- continue
- return {"knowledge_results": results}
- async def llm_answer(state: HealthButlerState) -> dict:
- """阶段3: LLM 生成个性化健康建议"""
- messages = [SystemMessage(content=HEALTH_BUTLER_PROMPT)]
- # 注入健康知识库上下文
- kb_ctx = state.get("kb_context", {})
- knowledge_results = state.get("knowledge_results", [])
- if knowledge_results:
- kb_text_parts = []
- for r in knowledge_results:
- d = r.get("data", {})
- name = d.get("itemName", r.get("keyword", ""))
- cat = d.get("category", "")
- desc = d.get("description", "")
- suggestion = d.get("suggestion", "")
- normal_range = d.get("normalRange", "")
- parts = [f"【{r['type']}】{name}"]
- if cat:
- parts.append(f"(分类: {cat})")
- if normal_range:
- parts.append(f"正常范围: {normal_range}")
- if desc:
- parts.append(f"{desc}")
- if suggestion:
- parts.append(f"建议: {suggestion}")
- kb_text_parts.append(" ".join(parts))
- kb_text = "\n---\n".join(kb_text_parts)
- messages.append(SystemMessage(content=f"健康知识库参考:\n{kb_text}"))
- # 注入家庭上下文
- family_id = state.get("family_id")
- if family_id:
- try:
- client = await java._get_client()
- resp = await client.post(
- "/api/ai/context",
- json={
- "user_id": str(state["user_id"]),
- "params": {
- "familyId": family_id,
- **({"childId": state["child_id"]} if state.get("child_id") else {}),
- },
- },
- )
- data = resp.json()
- if data.get("code") == 200:
- ctx = data.get("data", {})
- if ctx:
- messages.append(SystemMessage(content=f"家庭上下文:\n{json.dumps(ctx, ensure_ascii=False, indent=2)}"))
- except Exception as e:
- logger.debug("加载家庭上下文跳过: %s", e)
- # 注入长期记忆
- try:
- memories = await memory_mgr.recall(state["user_id"], state["query"])
- if memories:
- mem_text = "\n".join([f"- {m}" for m in memories])
- messages.append(SystemMessage(content=f"相关历史对话:\n{mem_text}"))
- except Exception as e:
- logger.warning("召回记忆失败: %s", e)
- # 用户消息
- focus = state.get("focus")
- query = state["query"]
- if focus:
- query = f"[重点关注: {focus}] {query}"
- messages.append(HumanMessage(content=query))
- response = await llm_with_tools.ainvoke(messages)
- answer = response.content
- # 提取任务
- import re
- tasks = []
- pattern = r'\[TASK:\s*\{([^}]+)\}\]'
- for match in re.findall(pattern, answer):
- task_info = {}
- for kv in match.split(","):
- if ":" in kv:
- k, v = kv.split(":", 1)
- task_info[k.strip()] = v.strip().strip('"').strip("'")
- if task_info.get("title"):
- tasks.append(task_info)
- sources = []
- for r in knowledge_results:
- d = r.get("data", {})
- sources.append({
- "type": "health_kb",
- "name": d.get("itemName", r.get("keyword", "")),
- "category": d.get("category", ""),
- })
- if response.response_metadata.get("tool_calls"):
- for tc in response.response_metadata["tool_calls"]:
- sources.append({
- "type": "tool",
- "name": tc.get("name", ""),
- })
- return {
- "answer": answer,
- "tasks": tasks,
- "sources": sources,
- "messages": [
- {"role": "user", "content": state["query"]},
- {"role": "assistant", "content": answer},
- ],
- }
- async def save_memory(state: HealthButlerState) -> dict:
- try:
- if state.get("messages"):
- await memory_mgr.save_conversation(
- state["user_id"],
- state["conversation_id"],
- state["messages"],
- )
- except Exception as e:
- logger.warning("保存记忆失败: %s", e)
- return {}
- # ── 路由 ──
- def route_after_knowledge(state: HealthButlerState) -> Literal["search_knowledge", "llm_answer"]:
- kb_ctx = state.get("kb_context", {})
- if kb_ctx.get("kb_loaded"):
- return "search_knowledge"
- return "llm_answer"
- # ── 构建图 ──
- builder.add_node("load_knowledge", load_knowledge)
- builder.add_node("search_knowledge", search_knowledge)
- builder.add_node("llm_answer", llm_answer)
- builder.add_node("save_memory", save_memory)
- builder.add_edge(START, "load_knowledge")
- builder.add_conditional_edges("load_knowledge", route_after_knowledge)
- builder.add_edge("search_knowledge", "llm_answer")
- builder.add_edge("llm_answer", "save_memory")
- builder.add_edge("save_memory", END)
- checkpointer = MemorySaver()
- graph = builder.compile(checkpointer=checkpointer)
- return graph
- def _extract_health_keywords(query: str) -> list[str]:
- """从 query 中提取健康相关关键词(菌属、指标、食物等名词)"""
- # 简单策略: 提取2字以上的中文词作为候选
- import re
- # 提取中文字符组成的连续词段
- chinese_words = re.findall(r'[\u4e00-\u9fff]{2,}', query)
- # 提取英文/拉丁名词组
- english_words = re.findall(r'[A-Za-z][A-Za-z\s]{2,20}[A-Za-z]', query)
- keywords = []
- # 过滤常见通用词
- stopwords = {"请问", "帮我", "能不能", "可不可以", "怎么样", "为什么",
- "是什么", "怎么办", "什么意思", "好不好", "是否", "你好",
- "谢谢", "请", "我想", "需要", "关于", "一个", "这个",
- "那个", "哪些", "什么", "怎么", "如何", "多少", "是不是"}
- for w in chinese_words:
- if w not in stopwords and len(w) >= 2:
- keywords.append(w)
- for w in english_words:
- w = w.strip()
- if len(w) > 2:
- keywords.append(w)
- # 去重 + 限数量
- seen = set()
- deduped = []
- for kw in keywords:
- if kw not in seen:
- seen.add(kw)
- deduped.append(kw)
- return deduped[:10]
|