| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372 |
- from __future__ import annotations
- """资源查询节点:调用ResourceSearchAgent执行航班与酒店搜索,汇总MCP工具返回结果。"""
- import json
- from typing import Any
- from langchain_core.messages import (
- AIMessage,
- ToolMessage,
- )
- from app.agents.resource_agent import (
- ResourceSearchAgent,
- )
- from app.graph.state import TravelState
- def extract_tool_payload(
- message: ToolMessage,
- ) -> dict[str, Any] | None:
- """从MCP ToolMessage中提取结构化结果。"""
- artifact = getattr(
- message,
- "artifact",
- None,
- )
- if isinstance(artifact, dict):
- structured = artifact.get(
- "structured_content"
- )
- if isinstance(structured, dict):
- return structured
- # 部分版本直接将结构放入artifact。
- if (
- "flights" in artifact
- or "hotels" in artifact
- or "suggestions" in artifact
- ):
- return artifact
- content = message.content
- if isinstance(content, str):
- try:
- parsed = json.loads(content)
- if isinstance(parsed, dict):
- return parsed
- except json.JSONDecodeError:
- return None
- if isinstance(content, list):
- for block in content:
- if not isinstance(block, dict):
- continue
- text = block.get("text")
- if not isinstance(text, str):
- continue
- try:
- parsed = json.loads(text)
- if isinstance(parsed, dict):
- return parsed
- except json.JSONDecodeError:
- continue
- return None
- def collect_tool_results(
- messages: list[Any],
- ) -> list[dict[str, Any]]:
- """收集Agent执行过的全部工具结果。"""
- tool_calls: dict[str, dict[str, Any]] = {}
- for message in messages:
- if not isinstance(message, AIMessage):
- continue
- for tool_call in message.tool_calls:
- tool_call_id = tool_call.get("id")
- if tool_call_id:
- tool_calls[tool_call_id] = {
- "name": tool_call.get("name"),
- "args": tool_call.get(
- "args",
- {},
- ),
- }
- records: list[dict[str, Any]] = []
- for message in messages:
- if not isinstance(message, ToolMessage):
- continue
- payload = extract_tool_payload(message)
- if payload is None:
- continue
- call_info = tool_calls.get(
- message.tool_call_id,
- {},
- )
- records.append(
- {
- "name": (
- message.name
- or call_info.get("name")
- ),
- "args": call_info.get(
- "args",
- {},
- ),
- "payload": payload,
- }
- )
- return records
- def extract_airport_codes(
- payload: dict[str, Any],
- ) -> list[str]:
- """从机场自动补全结果中提取城市机场代码。"""
- suggestions = payload.get(
- "suggestions",
- [],
- )
- if not isinstance(suggestions, list):
- return []
- city_suggestions = [
- suggestion
- for suggestion in suggestions
- if isinstance(suggestion, dict)
- and suggestion.get("type") == "city"
- ]
- candidates = (
- city_suggestions
- if city_suggestions
- else suggestions
- )
- for suggestion in candidates:
- if not isinstance(suggestion, dict):
- continue
- airports = suggestion.get(
- "airports",
- [],
- )
- if not isinstance(airports, list):
- continue
- codes: list[str] = []
- for airport in airports:
- if not isinstance(airport, dict):
- continue
- code = str(
- airport.get("code", "")
- ).strip().upper()
- if code and code not in codes:
- codes.append(code)
- if codes:
- return codes
- return []
- def make_resource_search_node(
- agent: ResourceSearchAgent,
- ):
- """创建LangGraph资源查询节点。"""
- async def resource_search_node(
- state: TravelState,
- ) -> dict[str, Any]:
- """资源查询节点:调用Agent并行搜索航班与酒店,将结构化结果写入state对应字段。"""
- request = state["travel_request"]
- try:
- agent_result = await agent.search(
- request
- )
- messages = agent_result.get(
- "messages",
- [],
- )
- records = collect_tool_results(
- messages
- )
- airport_records = [
- record
- for record in records
- if record["name"]
- == "search_airports"
- ]
- flight_records = [
- record
- for record in records
- if record["name"]
- == "search_flights"
- ]
- hotel_records = [
- record
- for record in records
- if record["name"]
- == "search_hotels"
- ]
- origin_payload = next(
- (
- record["payload"]
- for record in airport_records
- if record["payload"].get(
- "query"
- )
- == request.origin_city
- ),
- (
- airport_records[0]["payload"]
- if airport_records
- else {}
- ),
- )
- destination_payload = next(
- (
- record["payload"]
- for record in airport_records
- if record["payload"].get(
- "query"
- )
- == request.destination_city
- ),
- (
- airport_records[1]["payload"]
- if len(airport_records) > 1
- else {}
- ),
- )
- origin_airports = (
- extract_airport_codes(
- origin_payload
- )
- )
- destination_airports = (
- extract_airport_codes(
- destination_payload
- )
- )
- flight_result = (
- flight_records[-1]["payload"]
- if flight_records
- else {}
- )
- hotel_result = (
- hotel_records[-1]["payload"]
- if hotel_records
- else {}
- )
- errors: list[str] = []
- if not origin_airports:
- errors.append(
- "未解析出出发城市机场。"
- )
- if not destination_airports:
- errors.append(
- "未解析出目的城市机场。"
- )
- if not flight_result:
- errors.append(
- "资源Agent未返回航班结果。"
- )
- if not hotel_result:
- errors.append(
- "资源Agent未返回酒店结果。"
- )
- if errors:
- return {
- "resource_messages": messages,
- "errors": errors,
- "final_answer": (
- "资源查询失败:"
- + ";".join(errors)
- ),
- }
- flight_count = flight_result.get(
- "total_count",
- 0,
- )
- hotel_count = hotel_result.get(
- "total_count",
- 0,
- )
- return {
- "origin_airports": (
- origin_airports
- ),
- "destination_airports": (
- destination_airports
- ),
- "flight_search_result": (
- flight_result
- ),
- "hotel_search_result": (
- hotel_result
- ),
- "resource_messages": messages,
- "final_answer": (
- "真实旅行资源查询完成:"
- f"出发机场"
- f"{'、'.join(origin_airports)},"
- f"目的机场"
- f"{'、'.join(destination_airports)};"
- f"获得{flight_count}条航班候选,"
- f"{hotel_count}家酒店候选。"
- ),
- }
- except Exception as exc:
- return {
- "errors": [
- "资源查询Agent执行失败:"
- f"{type(exc).__name__}: {exc}"
- ],
- "final_answer": (
- "资源查询Agent执行失败:"
- f"{type(exc).__name__}: {exc}"
- ),
- }
- return resource_search_node
|