test_travel_mcp.py 8.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259
  1. from __future__ import annotations
  2. """直接调用自建旅行MCP服务器工具,验证返回结构是否符合预期。"""
  3. import asyncio
  4. from typing import Any
  5. from mcp import ClientSession
  6. from mcp.shared.memory import create_client_server_memory_streams
  7. from rich.console import Console
  8. from rich.table import Table
  9. from mcp_servers.travel_search_server import mcp
  10. console = Console()
  11. def get_payload(result: Any) -> dict[str, Any]:
  12. payload = result.structuredContent
  13. if not isinstance(payload, dict):
  14. raise RuntimeError(
  15. "MCP工具没有返回结构化结果。"
  16. )
  17. return payload
  18. async def main() -> None:
  19. console.rule(
  20. "[bold blue]TravelMind MCP Server测试"
  21. )
  22. # MCP SDK 1.x 内存直连模式:
  23. # 1. 创建双向内存流
  24. # 2. 服务端通过 asyncio.create_task 在后台运行
  25. # 3. 客户端通过 ClientSession 连接并初始化
  26. async with create_client_server_memory_streams() as (
  27. client_streams,
  28. server_streams,
  29. ):
  30. client_read, client_write = client_streams
  31. server_read, server_write = server_streams
  32. init_opts = mcp._mcp_server.create_initialization_options()
  33. async with ClientSession(client_read, client_write) as client:
  34. server_task = asyncio.create_task(
  35. mcp._mcp_server.run(
  36. server_read,
  37. server_write,
  38. init_opts,
  39. )
  40. )
  41. try:
  42. await client.initialize()
  43. tools_result = await client.list_tools()
  44. tool_table = Table(title="已发现MCP工具")
  45. tool_table.add_column("工具名")
  46. tool_table.add_column("说明")
  47. for tool in tools_result.tools:
  48. tool_table.add_row(
  49. tool.name,
  50. (tool.description or "")[:100],
  51. )
  52. console.print(tool_table)
  53. console.rule("[bold]1. 查询往返去程")
  54. outbound_result = await client.call_tool(
  55. "search_flights",
  56. {
  57. "departure_airports": [
  58. "SHA",
  59. "PVG",
  60. ],
  61. "arrival_airports": [
  62. "CTU",
  63. "TFU",
  64. ],
  65. "outbound_date": "2026-08-10",
  66. "return_date": "2026-08-13",
  67. "flight_type": "round_trip",
  68. "adults": 2,
  69. "currency": "CNY",
  70. },
  71. )
  72. outbound_payload = get_payload(
  73. outbound_result
  74. )
  75. outbound_flights = outbound_payload.get(
  76. "flights",
  77. [],
  78. )
  79. console.print(
  80. {
  81. "去程候选总数": (
  82. outbound_payload.get(
  83. "total_count"
  84. )
  85. ),
  86. "实际结构化数量": len(
  87. outbound_flights
  88. ),
  89. }
  90. )
  91. if not outbound_flights:
  92. raise RuntimeError(
  93. "没有查询到去程航班。"
  94. )
  95. selected_outbound = next(
  96. (
  97. flight
  98. for flight in outbound_flights
  99. if flight.get("departure_token")
  100. ),
  101. None,
  102. )
  103. if selected_outbound is None:
  104. raise RuntimeError(
  105. "去程候选中没有 departure_token。"
  106. )
  107. console.print(
  108. {
  109. "测试选中的去程": (
  110. selected_outbound.get(
  111. "flight_numbers"
  112. )
  113. ),
  114. "出发时间": (
  115. selected_outbound.get(
  116. "departure_time"
  117. )
  118. ),
  119. "到达时间": (
  120. selected_outbound.get(
  121. "arrival_time"
  122. )
  123. ),
  124. "价格": selected_outbound.get(
  125. "price"
  126. ),
  127. }
  128. )
  129. console.rule("[bold]2. 查询对应返程")
  130. return_result = await client.call_tool(
  131. "search_return_flights",
  132. {
  133. "departure_token": (
  134. selected_outbound[
  135. "departure_token"
  136. ]
  137. ),
  138. "departure_airports": [
  139. "SHA",
  140. "PVG",
  141. ],
  142. "arrival_airports": [
  143. "CTU",
  144. "TFU",
  145. ],
  146. "outbound_date": "2026-08-10",
  147. "return_date": "2026-08-13",
  148. "adults": 2,
  149. "currency": "CNY",
  150. },
  151. )
  152. return_payload = get_payload(
  153. return_result
  154. )
  155. console.print(
  156. {
  157. "返程候选总数": (
  158. return_payload.get(
  159. "total_count"
  160. )
  161. ),
  162. "实际结构化数量": len(
  163. return_payload.get(
  164. "flights",
  165. [],
  166. )
  167. ),
  168. }
  169. )
  170. console.rule("[bold]3. 查询酒店")
  171. hotel_result = await client.call_tool(
  172. "search_hotels",
  173. {
  174. "query": "成都酒店",
  175. "check_in_date": "2026-08-10",
  176. "check_out_date": "2026-08-13",
  177. "adults": 2,
  178. "currency": "CNY",
  179. },
  180. )
  181. hotel_payload = get_payload(
  182. hotel_result
  183. )
  184. hotels = hotel_payload.get("hotels", [])
  185. console.print(
  186. {
  187. "酒店候选总数": (
  188. hotel_payload.get(
  189. "total_count"
  190. )
  191. ),
  192. "实际结构化数量": len(hotels),
  193. }
  194. )
  195. if hotels:
  196. console.print(
  197. {
  198. "首家酒店": hotels[0].get(
  199. "name"
  200. ),
  201. "评分": hotels[0].get(
  202. "overall_rating"
  203. ),
  204. "每晚价格": hotels[0].get(
  205. "price_per_night"
  206. ),
  207. "总价格": hotels[0].get(
  208. "total_price"
  209. ),
  210. }
  211. )
  212. finally:
  213. server_task.cancel()
  214. try:
  215. await server_task
  216. except asyncio.CancelledError:
  217. pass
  218. console.rule(
  219. "[bold green]MCP Server测试通过"
  220. )
  221. if __name__ == "__main__":
  222. asyncio.run(main())