travel_search_server.py 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328
  1. from __future__ import annotations
  2. """自建旅行搜索MCP服务器:把TravelSearchService的航班/酒店搜索能力以MCP工具形式暴露,供LangChain Agent调用。"""
  3. from datetime import date, datetime
  4. from typing import Any
  5. from mcp.server.fastmcp import FastMCP
  6. from app.clients.serpapi_client import SerpApiClient, SerpApiError
  7. from app.config import get_settings
  8. from app.schemas.flight import FlightSearchQuery, FlightSearchResult
  9. from app.schemas.hotel import HotelSearchQuery, HotelSearchResult
  10. from app.services.travel_search_service import TravelSearchService
  11. mcp = FastMCP("travel-search")
  12. # ---------------------------------------------------------------------------
  13. # 辅助函数
  14. # ---------------------------------------------------------------------------
  15. def _parse_date(value: Any) -> date:
  16. """将字符串或 date 对象统一转为 date。"""
  17. if isinstance(value, date):
  18. return value
  19. if isinstance(value, str):
  20. return date.fromisoformat(value.strip())
  21. raise ValueError(f"无法解析日期: {value!r}")
  22. def _build_service() -> TravelSearchService:
  23. """创建 Service(懒加载 API key,避免模块导入时崩溃)。"""
  24. settings = get_settings()
  25. client = SerpApiClient(
  26. api_key=settings.require("serpapi_api_key"),
  27. timeout_seconds=settings.request_timeout_seconds,
  28. )
  29. return TravelSearchService(client)
  30. # ---------------------------------------------------------------------------
  31. # MCP 工具
  32. # ---------------------------------------------------------------------------
  33. @mcp.tool()
  34. async def search_flights(
  35. departure_airports: list[str],
  36. arrival_airports: list[str],
  37. outbound_date: str,
  38. return_date: str | None = None,
  39. flight_type: str = "round_trip",
  40. adults: int = 1,
  41. children: int = 0,
  42. travel_class: str = "economy",
  43. currency: str = "CNY",
  44. sort_by: str | None = None,
  45. stops: str | None = None,
  46. max_price: int | None = None,
  47. outbound_times: str | None = None,
  48. language: str = "zh-cn",
  49. country: str = "cn",
  50. no_cache: bool = False,
  51. show_hidden: bool = False,
  52. deep_search: bool = False,
  53. ) -> dict[str, Any]:
  54. """查询航班(单程或往返去程)。
  55. 返回标准化的航班候选列表。
  56. """
  57. query = FlightSearchQuery(
  58. departure_airports=departure_airports,
  59. arrival_airports=arrival_airports,
  60. outbound_date=_parse_date(outbound_date),
  61. return_date=_parse_date(return_date) if return_date else None,
  62. flight_type=flight_type, # type: ignore[arg-type]
  63. adults=adults,
  64. children=children,
  65. travel_class=travel_class, # type: ignore[arg-type]
  66. currency=currency,
  67. sort_by=sort_by, # type: ignore[arg-type]
  68. stops=stops, # type: ignore[arg-type]
  69. max_price=max_price,
  70. outbound_times=outbound_times,
  71. language=language,
  72. country=country,
  73. no_cache=no_cache,
  74. show_hidden=show_hidden,
  75. deep_search=deep_search,
  76. )
  77. service = _build_service()
  78. result: FlightSearchResult = await service.search_flights(query)
  79. return _serialize_flight_result(result)
  80. @mcp.tool()
  81. async def search_return_flights(
  82. departure_token: str,
  83. departure_airports: list[str],
  84. arrival_airports: list[str],
  85. outbound_date: str,
  86. return_date: str,
  87. adults: int = 1,
  88. children: int = 0,
  89. travel_class: str = "economy",
  90. currency: str = "CNY",
  91. sort_by: str | None = None,
  92. stops: str | None = None,
  93. max_price: int | None = None,
  94. language: str = "zh-cn",
  95. country: str = "cn",
  96. no_cache: bool = False,
  97. show_hidden: bool = False,
  98. deep_search: bool = False,
  99. ) -> dict[str, Any]:
  100. """查询与选定去程匹配的返程航班。
  101. 需要提供去程的 departure_token。
  102. """
  103. query = FlightSearchQuery(
  104. departure_airports=departure_airports,
  105. arrival_airports=arrival_airports,
  106. outbound_date=_parse_date(outbound_date),
  107. return_date=_parse_date(return_date),
  108. flight_type="round_trip",
  109. adults=adults,
  110. children=children,
  111. travel_class=travel_class, # type: ignore[arg-type]
  112. currency=currency,
  113. sort_by=sort_by, # type: ignore[arg-type]
  114. stops=stops, # type: ignore[arg-type]
  115. max_price=max_price,
  116. language=language,
  117. country=country,
  118. no_cache=no_cache,
  119. show_hidden=show_hidden,
  120. deep_search=deep_search,
  121. )
  122. service = _build_service()
  123. result: FlightSearchResult = await service.search_return_flights(
  124. query=query,
  125. departure_token=departure_token,
  126. )
  127. return _serialize_flight_result(result)
  128. @mcp.tool()
  129. async def search_hotels(
  130. query: str,
  131. check_in_date: str,
  132. check_out_date: str,
  133. adults: int = 1,
  134. children_ages: list[int] | None = None,
  135. currency: str = "CNY",
  136. language: str = "zh-cn",
  137. country: str = "cn",
  138. min_price: int | None = None,
  139. max_price: int | None = None,
  140. rating_filter: str | None = None,
  141. hotel_class: list[int] | None = None,
  142. free_cancellation: bool | None = None,
  143. sort: str | None = None,
  144. no_cache: bool = False,
  145. ) -> dict[str, Any]:
  146. """搜索酒店。
  147. 根据关键词和入住/退房日期返回标准化的酒店候选列表。
  148. """
  149. hotel_query = HotelSearchQuery(
  150. query=query,
  151. check_in_date=_parse_date(check_in_date),
  152. check_out_date=_parse_date(check_out_date),
  153. adults=adults,
  154. children_ages=children_ages or [],
  155. currency=currency,
  156. language=language,
  157. country=country,
  158. min_price=min_price,
  159. max_price=max_price,
  160. rating_filter=rating_filter, # type: ignore[arg-type]
  161. hotel_class=hotel_class or [], # type: ignore[arg-type]
  162. free_cancellation=free_cancellation,
  163. sort=sort, # type: ignore[arg-type]
  164. no_cache=no_cache,
  165. )
  166. service = _build_service()
  167. result: HotelSearchResult = await service.search_hotels(hotel_query)
  168. return _serialize_hotel_result(result)
  169. @mcp.tool()
  170. async def search_airports(
  171. query: str,
  172. ) -> dict[str, Any]:
  173. """根据城市或机场名称查询真实机场代码。
  174. 返回城市候选以及该城市包含的机场名称和
  175. IATA三字代码。查询航班前应先调用本工具。
  176. """
  177. service = _build_service()
  178. try:
  179. return await service.search_airports(query)
  180. except SerpApiError as exc:
  181. raise RuntimeError(
  182. f"机场查询失败:{exc}"
  183. ) from exc
  184. # ---------------------------------------------------------------------------
  185. # 序列化辅助:把 Pydantic 模型转成 JSON-safe dict
  186. # ---------------------------------------------------------------------------
  187. def _serialize_datetime(v: Any) -> str | None:
  188. """把datetime对象序列化为ISO格式字符串,None则返回None。"""
  189. if isinstance(v, datetime):
  190. return v.isoformat()
  191. return None
  192. def _serialize_flight_result(result: FlightSearchResult) -> dict[str, Any]:
  193. """把Pydantic航班结果模型序列化为MCP工具返回的JSON结构。"""
  194. flights: list[dict[str, Any]] = []
  195. for opt in result.flights:
  196. flights.append(
  197. {
  198. "option_id": opt.option_id,
  199. "source_group": opt.source_group,
  200. "provider_rank": opt.provider_rank,
  201. "flight_numbers": [
  202. seg.flight_number for seg in opt.segments
  203. ],
  204. "departure_airport_code": opt.departure_airport_code,
  205. "final_arrival_airport_code": opt.final_arrival_airport_code,
  206. "departure_time": _serialize_datetime(opt.departure_time),
  207. "arrival_time": _serialize_datetime(opt.arrival_time),
  208. "total_duration_minutes": opt.total_duration_minutes,
  209. "stop_count": opt.stop_count,
  210. "price": opt.price,
  211. "currency": opt.currency,
  212. "airlines": opt.airlines,
  213. "is_overnight": opt.is_overnight,
  214. "departure_token": opt.departure_token,
  215. "booking_token": opt.booking_token,
  216. "segments": [
  217. {
  218. "flight_number": seg.flight_number,
  219. "airline": seg.airline,
  220. "departure_airport_code": (
  221. seg.departure_airport_code
  222. ),
  223. "arrival_airport_code": (
  224. seg.arrival_airport_code
  225. ),
  226. "departure_time": _serialize_datetime(
  227. seg.departure_time
  228. ),
  229. "arrival_time": _serialize_datetime(
  230. seg.arrival_time
  231. ),
  232. "duration_minutes": seg.duration_minutes,
  233. }
  234. for seg in opt.segments
  235. ],
  236. "layovers": [
  237. {
  238. "airport_code": lay.airport_code,
  239. "duration_minutes": lay.duration_minutes,
  240. }
  241. for lay in opt.layovers
  242. ],
  243. }
  244. )
  245. return {
  246. "flights": flights,
  247. "total_count": result.total_count,
  248. "warnings": result.warnings,
  249. }
  250. def _serialize_hotel_result(
  251. result: HotelSearchResult,
  252. ) -> dict[str, Any]:
  253. """把Pydantic酒店结果模型序列化为MCP工具返回的JSON结构。"""
  254. hotels: list[dict[str, Any]] = []
  255. for h in result.hotels:
  256. hotels.append(
  257. {
  258. "hotel_id": h.hotel_id,
  259. "name": h.name,
  260. "description": h.description,
  261. "overall_rating": h.overall_rating,
  262. "review_count": h.review_count,
  263. "location_rating": h.location_rating,
  264. "price_per_night": h.price_per_night,
  265. "total_price": h.total_price,
  266. "currency": h.currency,
  267. "hotel_class": h.hotel_class,
  268. "amenities": h.amenities,
  269. "free_cancellation": h.free_cancellation,
  270. "thumbnail_url": h.thumbnail_url,
  271. "coordinates": (
  272. h.coordinates.model_dump()
  273. if h.coordinates is not None
  274. else None
  275. ),
  276. }
  277. )
  278. return {
  279. "hotels": hotels,
  280. "total_count": result.total_count,
  281. "warnings": result.warnings,
  282. }
  283. if __name__ == "__main__":
  284. mcp.run()