selection_builder.py 2.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129
  1. from __future__ import annotations
  2. """候选筛选阶段子图:对航班/酒店候选做确定性筛选与排名,为路线评估提供输入。"""
  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.agents.resource_agent import (
  12. ResourceSearchAgent,
  13. )
  14. from app.graph.nodes import (
  15. clarification_node,
  16. error_node,
  17. make_parse_request_node,
  18. route_after_parse,
  19. )
  20. from app.graph.resource_nodes import (
  21. make_resource_search_node,
  22. )
  23. from app.graph.selection_nodes import (
  24. make_candidate_selection_node,
  25. route_after_resources_to_selection,
  26. )
  27. from app.graph.state import TravelState
  28. from app.mcp_client import load_mcp_tools
  29. async def build_selection_graph():
  30. """构建资源查询与候选筛选工作流。"""
  31. tool_bundle = await load_mcp_tools()
  32. return_flight_tool = next(
  33. (
  34. tool
  35. for tool in tool_bundle.travel_tools
  36. if tool.name
  37. == "search_return_flights"
  38. ),
  39. None,
  40. )
  41. if return_flight_tool is None:
  42. raise RuntimeError(
  43. "没有加载到"
  44. "search_return_flights工具。"
  45. )
  46. requirement_agent = RequirementAgent()
  47. resource_agent = ResourceSearchAgent(
  48. tools=tool_bundle.travel_tools,
  49. )
  50. builder = StateGraph(TravelState)
  51. builder.add_node(
  52. "parse_request",
  53. make_parse_request_node(
  54. requirement_agent
  55. ),
  56. )
  57. builder.add_node(
  58. "clarify",
  59. clarification_node,
  60. )
  61. builder.add_node(
  62. "error",
  63. error_node,
  64. )
  65. builder.add_node(
  66. "search_resources",
  67. make_resource_search_node(
  68. resource_agent
  69. ),
  70. )
  71. builder.add_node(
  72. "select_candidates",
  73. make_candidate_selection_node(
  74. return_flight_tool,
  75. max_outbound_queries=2,
  76. max_returns_per_outbound=5,
  77. ),
  78. )
  79. builder.add_edge(
  80. START,
  81. "parse_request",
  82. )
  83. builder.add_conditional_edges(
  84. "parse_request",
  85. route_after_parse,
  86. {
  87. "clarify": "clarify",
  88. "error": "error",
  89. "ready": "search_resources",
  90. },
  91. )
  92. builder.add_edge("clarify", END)
  93. builder.add_edge("error", END)
  94. builder.add_conditional_edges(
  95. "search_resources",
  96. route_after_resources_to_selection,
  97. {
  98. "select_candidates": (
  99. "select_candidates"
  100. ),
  101. "end": END,
  102. },
  103. )
  104. builder.add_edge(
  105. "select_candidates",
  106. END,
  107. )
  108. return builder.compile()