from __future__ import annotations """自建旅行搜索MCP服务器:把TravelSearchService的航班/酒店搜索能力以MCP工具形式暴露,供LangChain Agent调用。""" from datetime import date, datetime from typing import Any from mcp.server.fastmcp import FastMCP from app.clients.serpapi_client import SerpApiClient, SerpApiError from app.config import get_settings from app.schemas.flight import FlightSearchQuery, FlightSearchResult from app.schemas.hotel import HotelSearchQuery, HotelSearchResult from app.services.travel_search_service import TravelSearchService mcp = FastMCP("travel-search") # --------------------------------------------------------------------------- # 辅助函数 # --------------------------------------------------------------------------- def _parse_date(value: Any) -> date: """将字符串或 date 对象统一转为 date。""" if isinstance(value, date): return value if isinstance(value, str): return date.fromisoformat(value.strip()) raise ValueError(f"无法解析日期: {value!r}") def _build_service() -> TravelSearchService: """创建 Service(懒加载 API key,避免模块导入时崩溃)。""" settings = get_settings() client = SerpApiClient( api_key=settings.require("serpapi_api_key"), timeout_seconds=settings.request_timeout_seconds, ) return TravelSearchService(client) # --------------------------------------------------------------------------- # MCP 工具 # --------------------------------------------------------------------------- @mcp.tool() async def search_flights( departure_airports: list[str], arrival_airports: list[str], outbound_date: str, return_date: str | None = None, flight_type: str = "round_trip", adults: int = 1, children: int = 0, travel_class: str = "economy", currency: str = "CNY", sort_by: str | None = None, stops: str | None = None, max_price: int | None = None, outbound_times: str | None = None, language: str = "zh-cn", country: str = "cn", no_cache: bool = False, show_hidden: bool = False, deep_search: bool = False, ) -> dict[str, Any]: """查询航班(单程或往返去程)。 返回标准化的航班候选列表。 """ query = FlightSearchQuery( departure_airports=departure_airports, arrival_airports=arrival_airports, outbound_date=_parse_date(outbound_date), return_date=_parse_date(return_date) if return_date else None, flight_type=flight_type, # type: ignore[arg-type] adults=adults, children=children, travel_class=travel_class, # type: ignore[arg-type] currency=currency, sort_by=sort_by, # type: ignore[arg-type] stops=stops, # type: ignore[arg-type] max_price=max_price, outbound_times=outbound_times, language=language, country=country, no_cache=no_cache, show_hidden=show_hidden, deep_search=deep_search, ) service = _build_service() result: FlightSearchResult = await service.search_flights(query) return _serialize_flight_result(result) @mcp.tool() async def search_return_flights( departure_token: str, departure_airports: list[str], arrival_airports: list[str], outbound_date: str, return_date: str, adults: int = 1, children: int = 0, travel_class: str = "economy", currency: str = "CNY", sort_by: str | None = None, stops: str | None = None, max_price: int | None = None, language: str = "zh-cn", country: str = "cn", no_cache: bool = False, show_hidden: bool = False, deep_search: bool = False, ) -> dict[str, Any]: """查询与选定去程匹配的返程航班。 需要提供去程的 departure_token。 """ query = FlightSearchQuery( departure_airports=departure_airports, arrival_airports=arrival_airports, outbound_date=_parse_date(outbound_date), return_date=_parse_date(return_date), flight_type="round_trip", adults=adults, children=children, travel_class=travel_class, # type: ignore[arg-type] currency=currency, sort_by=sort_by, # type: ignore[arg-type] stops=stops, # type: ignore[arg-type] max_price=max_price, language=language, country=country, no_cache=no_cache, show_hidden=show_hidden, deep_search=deep_search, ) service = _build_service() result: FlightSearchResult = await service.search_return_flights( query=query, departure_token=departure_token, ) return _serialize_flight_result(result) @mcp.tool() async def search_hotels( query: str, check_in_date: str, check_out_date: str, adults: int = 1, children_ages: list[int] | None = None, currency: str = "CNY", language: str = "zh-cn", country: str = "cn", min_price: int | None = None, max_price: int | None = None, rating_filter: str | None = None, hotel_class: list[int] | None = None, free_cancellation: bool | None = None, sort: str | None = None, no_cache: bool = False, ) -> dict[str, Any]: """搜索酒店。 根据关键词和入住/退房日期返回标准化的酒店候选列表。 """ hotel_query = HotelSearchQuery( query=query, check_in_date=_parse_date(check_in_date), check_out_date=_parse_date(check_out_date), adults=adults, children_ages=children_ages or [], currency=currency, language=language, country=country, min_price=min_price, max_price=max_price, rating_filter=rating_filter, # type: ignore[arg-type] hotel_class=hotel_class or [], # type: ignore[arg-type] free_cancellation=free_cancellation, sort=sort, # type: ignore[arg-type] no_cache=no_cache, ) service = _build_service() result: HotelSearchResult = await service.search_hotels(hotel_query) return _serialize_hotel_result(result) @mcp.tool() async def search_airports( query: str, ) -> dict[str, Any]: """根据城市或机场名称查询真实机场代码。 返回城市候选以及该城市包含的机场名称和 IATA三字代码。查询航班前应先调用本工具。 """ service = _build_service() try: return await service.search_airports(query) except SerpApiError as exc: raise RuntimeError( f"机场查询失败:{exc}" ) from exc # --------------------------------------------------------------------------- # 序列化辅助:把 Pydantic 模型转成 JSON-safe dict # --------------------------------------------------------------------------- def _serialize_datetime(v: Any) -> str | None: """把datetime对象序列化为ISO格式字符串,None则返回None。""" if isinstance(v, datetime): return v.isoformat() return None def _serialize_flight_result(result: FlightSearchResult) -> dict[str, Any]: """把Pydantic航班结果模型序列化为MCP工具返回的JSON结构。""" flights: list[dict[str, Any]] = [] for opt in result.flights: flights.append( { "option_id": opt.option_id, "source_group": opt.source_group, "provider_rank": opt.provider_rank, "flight_numbers": [ seg.flight_number for seg in opt.segments ], "departure_airport_code": opt.departure_airport_code, "final_arrival_airport_code": opt.final_arrival_airport_code, "departure_time": _serialize_datetime(opt.departure_time), "arrival_time": _serialize_datetime(opt.arrival_time), "total_duration_minutes": opt.total_duration_minutes, "stop_count": opt.stop_count, "price": opt.price, "currency": opt.currency, "airlines": opt.airlines, "is_overnight": opt.is_overnight, "departure_token": opt.departure_token, "booking_token": opt.booking_token, "segments": [ { "flight_number": seg.flight_number, "airline": seg.airline, "departure_airport_code": ( seg.departure_airport_code ), "arrival_airport_code": ( seg.arrival_airport_code ), "departure_time": _serialize_datetime( seg.departure_time ), "arrival_time": _serialize_datetime( seg.arrival_time ), "duration_minutes": seg.duration_minutes, } for seg in opt.segments ], "layovers": [ { "airport_code": lay.airport_code, "duration_minutes": lay.duration_minutes, } for lay in opt.layovers ], } ) return { "flights": flights, "total_count": result.total_count, "warnings": result.warnings, } def _serialize_hotel_result( result: HotelSearchResult, ) -> dict[str, Any]: """把Pydantic酒店结果模型序列化为MCP工具返回的JSON结构。""" hotels: list[dict[str, Any]] = [] for h in result.hotels: hotels.append( { "hotel_id": h.hotel_id, "name": h.name, "description": h.description, "overall_rating": h.overall_rating, "review_count": h.review_count, "location_rating": h.location_rating, "price_per_night": h.price_per_night, "total_price": h.total_price, "currency": h.currency, "hotel_class": h.hotel_class, "amenities": h.amenities, "free_cancellation": h.free_cancellation, "thumbnail_url": h.thumbnail_url, "coordinates": ( h.coordinates.model_dump() if h.coordinates is not None else None ), } ) return { "hotels": hotels, "total_count": result.total_count, "warnings": result.warnings, } if __name__ == "__main__": mcp.run()