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]