from typing import TypedDict, Literal from langgraph.graph import StateGraph, START, END from langgraph.checkpoint.memory import MemorySaver from langchain_openai import ChatOpenAI from langchain_core.messages import SystemMessage, HumanMessage from app.agents.intent_classifier import IntentClassifier, Intent from app.agents.chat_agent import ChatAgent from app.tools.product_tools import ( search_product_by_keyword, search_article_by_keyword, search_activity_by_keyword, ) from app.memory.store import MemoryManager from app.config import settings from app.prompt_service import get_prompt from app.portrait_service import get_user_portrait_text from app.tools.java_client import JavaClient import logging logger = logging.getLogger(__name__) class ChatState(TypedDict): query: str user_id: int conversation_id: str child_id: int | None intent: Intent | None context: dict | None self_check_result: str | None # P1-2: 自检结果 JSON messages: list | None answer: str | None tasks: list[dict] sources: list[dict] DEFAULT_CHAT_PROMPT = """你是一个儿童成长家庭助手, 回答关于孩子成长、健康、教育的各种问题。 你可以使用搜索工具查找商品、活动和文章来辅助回答。 回答原则: 1. 用中文, 语气温暖亲切 2. 如果用户提到具体孩子, 参考提供的家庭上下文 3. 需要推荐时使用搜索工具 4. 可以生成 [TASK: {"title": "任务名", "dimension": "身/心/智/行/富", "points": 10}] 标记来创建行动任务 5. 不要编造医疗建议, 严重问题建议咨询医生 """ def create_chat_graph(): """创建聊天 StateGraph""" agent = ChatAgent() classifier = IntentClassifier() llm = ChatOpenAI( model=settings.llm_model, api_key=settings.llm_api_key, base_url=settings.llm_base_url, temperature=settings.llm_temperature, ) # 初始化嵌入模型用于记忆存储 from app.rag.embeddings import get_embeddings embeddings = get_embeddings() memory_mgr = MemoryManager(embeddings) java = JavaClient() llm_with_tools = llm.bind_tools([ search_product_by_keyword, search_article_by_keyword, search_activity_by_keyword, ]) builder = StateGraph(ChatState) # ── 节点 ── async def classify_intent(state: ChatState) -> dict: intent = await classifier.classify( state["query"], context=str(state.get("context", {})), ) return {"intent": intent} async def load_context(state: ChatState) -> dict: ctx = await agent.load_context(state["user_id"], state.get("child_id")) return {"context": ctx} async def llm_call(state: ChatState) -> dict: """核心 LLM 调用 + Tool""" chat_prompt = await get_prompt("chat_assistant") or DEFAULT_CHAT_PROMPT messages = [SystemMessage(content=chat_prompt)] # 画像注入 member_id = state.get("child_id") portrait_text = await get_user_portrait_text(java, state["user_id"], member_id) if portrait_text: messages.insert(1, SystemMessage(content=portrait_text)) # P1-2: 自检结果注入(如有) self_check_result = state.get("self_check_result") if self_check_result: try: import json as _json sc = _json.loads(self_check_result) sc_text = f"用户最近一次五维自检结果:\n总分 {sc.get('totalScore', 0)}/45," dims = sc.get('dimensions') or sc.get('scores') or {} if isinstance(dims, list): dim_lines = [f"{d.get('name','')}: {d.get('score',0)}分" for d in dims] elif isinstance(dims, dict): dim_lines = [f"{k}: {v}分" for k, v in dims.items()] else: dim_lines = [str(dims)] sc_text += " | ".join(dim_lines) insight = sc.get('familyInsight') or sc.get('aiInsight') if insight: sc_text += f"\nAI 洞察: {insight}" messages.insert(1, SystemMessage(content=sc_text)) except Exception as _e: logger.warning("解析 self_check_result 失败: %s", _e) # 注入家庭上下文 + 用户健康现状 ctx = state.get("context", {}) if isinstance(ctx, dict): ctx_parts = [] if ctx.get("child_id") is not None: ctx_parts.append(f"child_id: {ctx['child_id']}") if ctx.get("report_id") is not None: ctx_parts.append(f"report_id: {ctx['report_id']}") if ctx.get("family_id") is not None: ctx_parts.append(f"family_id: {ctx['family_id']}") if ctx.get("health_status"): ctx_parts.append(f"用户健康现状: {ctx['health_status']}") if ctx.get("diet_preferences"): ctx_parts.append(f"用户饮食偏好: {ctx['diet_preferences']}") if ctx_parts: messages.append(SystemMessage(content="\n家庭上下文:\n" + "\n".join(ctx_parts))) # 注入长期记忆 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) # 用户消息 messages.append(HumanMessage(content=state["query"])) response = await llm_with_tools.ainvoke(messages) answer = response.content # 提取任务 tasks = await agent.extract_tasks(answer, state["user_id"], state["conversation_id"]) # 提取来源 sources = [] if response.response_metadata.get("tool_calls"): for tc in response.response_metadata["tool_calls"]: sources.append({ "type": "tool", "name": tc.get("name", ""), "input": tc.get("args", {}), }) return { "answer": answer, "tasks": tasks, "sources": sources, "messages": [{"role": "user", "content": state["query"]}, {"role": "assistant", "content": answer}], } async def save_memory(state: ChatState) -> 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_by_intent(state: ChatState) -> Literal["llm_call", END]: if state["intent"] in ( Intent.RECOMMEND, Intent.ANALYSIS, Intent.HEALTH, ): # 这些意图需要更专业的 Agent (Phase 3 实现) # 当前先走通用 LLM pass return "llm_call" # ── 构建图 ── builder.add_node("classify_intent", classify_intent) builder.add_node("load_context", load_context) builder.add_node("llm_call", llm_call) builder.add_node("save_memory", save_memory) builder.add_edge(START, "classify_intent") builder.add_edge("classify_intent", "load_context") builder.add_conditional_edges("load_context", route_by_intent) builder.add_edge("llm_call", "save_memory") builder.add_edge("save_memory", END) checkpointer = MemorySaver() graph = builder.compile(checkpointer=checkpointer) return graph