resource_nodes.py 9.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372
  1. from __future__ import annotations
  2. """资源查询节点:调用ResourceSearchAgent执行航班与酒店搜索,汇总MCP工具返回结果。"""
  3. import json
  4. from typing import Any
  5. from langchain_core.messages import (
  6. AIMessage,
  7. ToolMessage,
  8. )
  9. from app.agents.resource_agent import (
  10. ResourceSearchAgent,
  11. )
  12. from app.graph.state import TravelState
  13. def extract_tool_payload(
  14. message: ToolMessage,
  15. ) -> dict[str, Any] | None:
  16. """从MCP ToolMessage中提取结构化结果。"""
  17. artifact = getattr(
  18. message,
  19. "artifact",
  20. None,
  21. )
  22. if isinstance(artifact, dict):
  23. structured = artifact.get(
  24. "structured_content"
  25. )
  26. if isinstance(structured, dict):
  27. return structured
  28. # 部分版本直接将结构放入artifact。
  29. if (
  30. "flights" in artifact
  31. or "hotels" in artifact
  32. or "suggestions" in artifact
  33. ):
  34. return artifact
  35. content = message.content
  36. if isinstance(content, str):
  37. try:
  38. parsed = json.loads(content)
  39. if isinstance(parsed, dict):
  40. return parsed
  41. except json.JSONDecodeError:
  42. return None
  43. if isinstance(content, list):
  44. for block in content:
  45. if not isinstance(block, dict):
  46. continue
  47. text = block.get("text")
  48. if not isinstance(text, str):
  49. continue
  50. try:
  51. parsed = json.loads(text)
  52. if isinstance(parsed, dict):
  53. return parsed
  54. except json.JSONDecodeError:
  55. continue
  56. return None
  57. def collect_tool_results(
  58. messages: list[Any],
  59. ) -> list[dict[str, Any]]:
  60. """收集Agent执行过的全部工具结果。"""
  61. tool_calls: dict[str, dict[str, Any]] = {}
  62. for message in messages:
  63. if not isinstance(message, AIMessage):
  64. continue
  65. for tool_call in message.tool_calls:
  66. tool_call_id = tool_call.get("id")
  67. if tool_call_id:
  68. tool_calls[tool_call_id] = {
  69. "name": tool_call.get("name"),
  70. "args": tool_call.get(
  71. "args",
  72. {},
  73. ),
  74. }
  75. records: list[dict[str, Any]] = []
  76. for message in messages:
  77. if not isinstance(message, ToolMessage):
  78. continue
  79. payload = extract_tool_payload(message)
  80. if payload is None:
  81. continue
  82. call_info = tool_calls.get(
  83. message.tool_call_id,
  84. {},
  85. )
  86. records.append(
  87. {
  88. "name": (
  89. message.name
  90. or call_info.get("name")
  91. ),
  92. "args": call_info.get(
  93. "args",
  94. {},
  95. ),
  96. "payload": payload,
  97. }
  98. )
  99. return records
  100. def extract_airport_codes(
  101. payload: dict[str, Any],
  102. ) -> list[str]:
  103. """从机场自动补全结果中提取城市机场代码。"""
  104. suggestions = payload.get(
  105. "suggestions",
  106. [],
  107. )
  108. if not isinstance(suggestions, list):
  109. return []
  110. city_suggestions = [
  111. suggestion
  112. for suggestion in suggestions
  113. if isinstance(suggestion, dict)
  114. and suggestion.get("type") == "city"
  115. ]
  116. candidates = (
  117. city_suggestions
  118. if city_suggestions
  119. else suggestions
  120. )
  121. for suggestion in candidates:
  122. if not isinstance(suggestion, dict):
  123. continue
  124. airports = suggestion.get(
  125. "airports",
  126. [],
  127. )
  128. if not isinstance(airports, list):
  129. continue
  130. codes: list[str] = []
  131. for airport in airports:
  132. if not isinstance(airport, dict):
  133. continue
  134. code = str(
  135. airport.get("code", "")
  136. ).strip().upper()
  137. if code and code not in codes:
  138. codes.append(code)
  139. if codes:
  140. return codes
  141. return []
  142. def make_resource_search_node(
  143. agent: ResourceSearchAgent,
  144. ):
  145. """创建LangGraph资源查询节点。"""
  146. async def resource_search_node(
  147. state: TravelState,
  148. ) -> dict[str, Any]:
  149. """资源查询节点:调用Agent并行搜索航班与酒店,将结构化结果写入state对应字段。"""
  150. request = state["travel_request"]
  151. try:
  152. agent_result = await agent.search(
  153. request
  154. )
  155. messages = agent_result.get(
  156. "messages",
  157. [],
  158. )
  159. records = collect_tool_results(
  160. messages
  161. )
  162. airport_records = [
  163. record
  164. for record in records
  165. if record["name"]
  166. == "search_airports"
  167. ]
  168. flight_records = [
  169. record
  170. for record in records
  171. if record["name"]
  172. == "search_flights"
  173. ]
  174. hotel_records = [
  175. record
  176. for record in records
  177. if record["name"]
  178. == "search_hotels"
  179. ]
  180. origin_payload = next(
  181. (
  182. record["payload"]
  183. for record in airport_records
  184. if record["payload"].get(
  185. "query"
  186. )
  187. == request.origin_city
  188. ),
  189. (
  190. airport_records[0]["payload"]
  191. if airport_records
  192. else {}
  193. ),
  194. )
  195. destination_payload = next(
  196. (
  197. record["payload"]
  198. for record in airport_records
  199. if record["payload"].get(
  200. "query"
  201. )
  202. == request.destination_city
  203. ),
  204. (
  205. airport_records[1]["payload"]
  206. if len(airport_records) > 1
  207. else {}
  208. ),
  209. )
  210. origin_airports = (
  211. extract_airport_codes(
  212. origin_payload
  213. )
  214. )
  215. destination_airports = (
  216. extract_airport_codes(
  217. destination_payload
  218. )
  219. )
  220. flight_result = (
  221. flight_records[-1]["payload"]
  222. if flight_records
  223. else {}
  224. )
  225. hotel_result = (
  226. hotel_records[-1]["payload"]
  227. if hotel_records
  228. else {}
  229. )
  230. errors: list[str] = []
  231. if not origin_airports:
  232. errors.append(
  233. "未解析出出发城市机场。"
  234. )
  235. if not destination_airports:
  236. errors.append(
  237. "未解析出目的城市机场。"
  238. )
  239. if not flight_result:
  240. errors.append(
  241. "资源Agent未返回航班结果。"
  242. )
  243. if not hotel_result:
  244. errors.append(
  245. "资源Agent未返回酒店结果。"
  246. )
  247. if errors:
  248. return {
  249. "resource_messages": messages,
  250. "errors": errors,
  251. "final_answer": (
  252. "资源查询失败:"
  253. + ";".join(errors)
  254. ),
  255. }
  256. flight_count = flight_result.get(
  257. "total_count",
  258. 0,
  259. )
  260. hotel_count = hotel_result.get(
  261. "total_count",
  262. 0,
  263. )
  264. return {
  265. "origin_airports": (
  266. origin_airports
  267. ),
  268. "destination_airports": (
  269. destination_airports
  270. ),
  271. "flight_search_result": (
  272. flight_result
  273. ),
  274. "hotel_search_result": (
  275. hotel_result
  276. ),
  277. "resource_messages": messages,
  278. "final_answer": (
  279. "真实旅行资源查询完成:"
  280. f"出发机场"
  281. f"{'、'.join(origin_airports)},"
  282. f"目的机场"
  283. f"{'、'.join(destination_airports)};"
  284. f"获得{flight_count}条航班候选,"
  285. f"{hotel_count}家酒店候选。"
  286. ),
  287. }
  288. except Exception as exc:
  289. return {
  290. "errors": [
  291. "资源查询Agent执行失败:"
  292. f"{type(exc).__name__}: {exc}"
  293. ],
  294. "final_answer": (
  295. "资源查询Agent执行失败:"
  296. f"{type(exc).__name__}: {exc}"
  297. ),
  298. }
  299. return resource_search_node