chat_graph.py 5.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169
  1. from typing import TypedDict, Literal
  2. from langgraph.graph import StateGraph, START, END
  3. from langgraph.checkpoint import MemorySaver
  4. from langchain_openai import ChatOpenAI
  5. from langchain_core.messages import SystemMessage, HumanMessage
  6. from app.agents.intent_classifier import IntentClassifier, Intent
  7. from app.agents.chat_agent import ChatAgent
  8. from app.tools.product_tools import (
  9. search_product_by_keyword,
  10. search_article_by_keyword,
  11. search_activity_by_keyword,
  12. )
  13. from app.memory.store import MemoryManager
  14. from app.config import settings
  15. import logging
  16. logger = logging.getLogger(__name__)
  17. class ChatState(TypedDict):
  18. query: str
  19. user_id: int
  20. conversation_id: str
  21. child_id: int | None
  22. intent: Intent | None
  23. context: dict | None
  24. messages: list | None
  25. answer: str | None
  26. tasks: list[dict]
  27. sources: list[dict]
  28. CHAT_SYSTEM_PROMPT = """你是一个儿童成长家庭助手, 回答关于孩子成长、健康、教育的各种问题。
  29. 你可以使用搜索工具查找商品、活动和文章来辅助回答。
  30. 回答原则:
  31. 1. 用中文, 语气温暖亲切
  32. 2. 如果用户提到具体孩子, 参考提供的家庭上下文
  33. 3. 需要推荐时使用搜索工具
  34. 4. 可以生成 [TASK: {"title": "任务名", "dimension": "身/心/智/行/富", "points": 10}] 标记来创建行动任务
  35. 5. 不要编造医疗建议, 严重问题建议咨询医生
  36. """
  37. def create_chat_graph():
  38. """创建聊天 StateGraph"""
  39. agent = ChatAgent()
  40. classifier = IntentClassifier()
  41. memory_mgr = MemoryManager()
  42. llm = ChatOpenAI(
  43. model=settings.llm_model,
  44. api_key=settings.llm_api_key,
  45. base_url=settings.llm_base_url,
  46. temperature=0.7,
  47. )
  48. llm_with_tools = llm.bind_tools([
  49. search_product_by_keyword,
  50. search_article_by_keyword,
  51. search_activity_by_keyword,
  52. ])
  53. builder = StateGraph(ChatState)
  54. # ── 节点 ──
  55. async def classify_intent(state: ChatState) -> dict:
  56. intent = await classifier.classify(
  57. state["query"],
  58. context=str(state.get("context", {})),
  59. )
  60. return {"intent": intent}
  61. async def load_context(state: ChatState) -> dict:
  62. ctx = await agent.load_context(state["user_id"], state.get("child_id"))
  63. return {"context": ctx}
  64. async def llm_call(state: ChatState) -> dict:
  65. """核心 LLM 调用 + Tool"""
  66. messages = [SystemMessage(content=CHAT_SYSTEM_PROMPT)]
  67. # 注入家庭上下文
  68. ctx = state.get("context", {})
  69. if ctx:
  70. ctx_text = f"\n家庭上下文:\n{ctx}"
  71. messages.append(SystemMessage(content=ctx_text))
  72. # 注入长期记忆
  73. try:
  74. memories = await memory_mgr.recall(state["user_id"], state["query"])
  75. if memories:
  76. mem_text = "\n".join([f"- {m}" for m in memories])
  77. messages.append(SystemMessage(
  78. content=f"相关历史对话:\n{mem_text}"
  79. ))
  80. except Exception as e:
  81. logger.warning("召回记忆失败: %s", e)
  82. # 用户消息
  83. messages.append(HumanMessage(content=state["query"]))
  84. response = await llm_with_tools.ainvoke(messages)
  85. answer = response.content
  86. # 提取任务
  87. tasks = await agent.extract_tasks(answer, state["user_id"], state["conversation_id"])
  88. # 提取来源
  89. sources = []
  90. if response.response_metadata.get("tool_calls"):
  91. for tc in response.response_metadata["tool_calls"]:
  92. sources.append({
  93. "type": "tool",
  94. "name": tc.get("name", ""),
  95. "input": tc.get("args", {}),
  96. })
  97. return {
  98. "answer": answer,
  99. "tasks": tasks,
  100. "sources": sources,
  101. "messages": [{"role": "user", "content": state["query"]},
  102. {"role": "assistant", "content": answer}],
  103. }
  104. async def save_memory(state: ChatState) -> dict:
  105. """对话后保存到长期记忆"""
  106. try:
  107. if state.get("messages"):
  108. await memory_mgr.save_conversation(
  109. state["user_id"],
  110. state["conversation_id"],
  111. state["messages"],
  112. )
  113. except Exception as e:
  114. logger.warning("保存记忆失败: %s", e)
  115. return {}
  116. # ── 路由 ──
  117. def route_by_intent(state: ChatState) -> Literal["llm_call", END]:
  118. if state["intent"] in (
  119. Intent.RECOMMEND,
  120. Intent.ANALYSIS,
  121. Intent.HEALTH,
  122. ):
  123. # 这些意图需要更专业的 Agent (Phase 3 实现)
  124. # 当前先走通用 LLM
  125. pass
  126. return "llm_call"
  127. # ── 构建图 ──
  128. builder.add_node("classify_intent", classify_intent)
  129. builder.add_node("load_context", load_context)
  130. builder.add_node("llm_call", llm_call)
  131. builder.add_node("save_memory", save_memory)
  132. builder.add_edge(START, "classify_intent")
  133. builder.add_edge("classify_intent", "load_context")
  134. builder.add_conditional_edges("load_context", route_by_intent)
  135. builder.add_edge("llm_call", "save_memory")
  136. builder.add_edge("save_memory", END)
  137. checkpointer = MemorySaver()
  138. graph = builder.compile(checkpointer=checkpointer)
  139. return graph