builder.py 1.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475
  1. from __future__ import annotations
  2. """需求解析阶段子图:从用户原始输入解析为结构化TravelRequest,路由到澄清/就绪/错误节点。"""
  3. from langgraph.graph import (
  4. END,
  5. START,
  6. StateGraph,
  7. )
  8. from app.agents.requirement_agent import (
  9. RequirementAgent,
  10. )
  11. from app.graph.nodes import (
  12. clarification_node,
  13. error_node,
  14. make_parse_request_node,
  15. ready_node,
  16. route_after_parse,
  17. )
  18. from app.graph.state import TravelState
  19. def build_requirement_graph(
  20. agent: RequirementAgent | None = None,
  21. ):
  22. """构建需求解析阶段的LangGraph。"""
  23. requirement_agent = (
  24. agent or RequirementAgent()
  25. )
  26. builder = StateGraph(TravelState)
  27. builder.add_node(
  28. "parse_request",
  29. make_parse_request_node(
  30. requirement_agent
  31. ),
  32. )
  33. builder.add_node(
  34. "clarify",
  35. clarification_node,
  36. )
  37. builder.add_node(
  38. "ready",
  39. ready_node,
  40. )
  41. builder.add_node(
  42. "error",
  43. error_node,
  44. )
  45. builder.add_edge(
  46. START,
  47. "parse_request",
  48. )
  49. builder.add_conditional_edges(
  50. "parse_request",
  51. route_after_parse,
  52. {
  53. "clarify": "clarify",
  54. "ready": "ready",
  55. "error": "error",
  56. },
  57. )
  58. builder.add_edge("clarify", END)
  59. builder.add_edge("ready", END)
  60. builder.add_edge("error", END)
  61. return builder.compile()