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