| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328 |
- 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()
|