test_mcp_protocol.py 14 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397
  1. import io
  2. import json
  3. import unittest
  4. from app import GatewayApp
  5. from services.token_store import InMemoryTokenStore
  6. from mcp_protocol import McpProtocolHandler
  7. class DummyApiClient:
  8. def __init__(self):
  9. self.calls = []
  10. def list_enabled_tools(self):
  11. return {
  12. 'code': 'MCP_0000',
  13. 'data': {
  14. 'tool_codes': [
  15. 'query_order',
  16. 'query_track',
  17. 'query_order_exact',
  18. 'list_order_filter_options',
  19. ],
  20. },
  21. }
  22. def call_tool(self, tool_code, route_path, payload, request_id):
  23. self.calls.append(
  24. {
  25. 'tool_code': tool_code,
  26. 'route_path': route_path,
  27. 'payload': payload,
  28. 'request_id': request_id,
  29. }
  30. )
  31. return {
  32. 'code': 'MCP_0000',
  33. 'msg': 'success',
  34. 'data': {
  35. 'summary': 'matched 1 order',
  36. 'records': [
  37. {
  38. 'order_no': 'SO20260706001',
  39. }
  40. ],
  41. 'tips': ['scoped by employee permissions'],
  42. },
  43. 'meta': {
  44. 'request_id': request_id,
  45. },
  46. }
  47. class BusinessErrorApiClient(DummyApiClient):
  48. def call_tool(self, tool_code, route_path, payload, request_id):
  49. super().call_tool(tool_code, route_path, payload, request_id)
  50. return {
  51. 'code': 'MCP_1301',
  52. 'msg': 'no order query permission',
  53. 'data': [],
  54. 'meta': {
  55. 'request_id': request_id,
  56. },
  57. }
  58. class FullColumnsApiClient(DummyApiClient):
  59. def call_tool(self, tool_code, route_path, payload, request_id):
  60. response = super().call_tool(tool_code, route_path, payload, request_id)
  61. columns = [
  62. ('order_number', '订单号'),
  63. ('reference_number', '客户参考号'),
  64. ('status_txt_name', '状态'),
  65. ('check_status_txt_name', '是否已查验'),
  66. ('customer_name', '客户名称'),
  67. ('customer_account_type_name', '客户属性'),
  68. ('inbound_date', '入库时间'),
  69. ('wo_num', '未完成工单'),
  70. ('product_name', '物流产品'),
  71. ('inbound_pieces', '件数'),
  72. ('inbound_volume', '体积(CBM)'),
  73. ('inbound_weight', '重量(KG)'),
  74. ('pro_cn_name', '品名'),
  75. ('export_declaration_type', '报关方式'),
  76. ('merge_declare_number', '合并报关单号'),
  77. ('delivery_address', '派送地址'),
  78. ('container_code', '柜号'),
  79. ('out_status_txt', '排舱单状态'),
  80. ('hinge_of_destination', '目的港'),
  81. ('etd', 'ETD'),
  82. ('atd', 'ATD'),
  83. ('eta', 'ETA'),
  84. ('ata', 'ATA'),
  85. ('release_time', '清关放行时间'),
  86. ('oversea_inbound_date', '海外入库时间'),
  87. ('appt_time', 'APPT时间'),
  88. ('est_loading_time', '预计装柜时间'),
  89. ('pickup_time', '海外提柜时间'),
  90. ('delivery_way_title', '派送方式'),
  91. ('tracking_number', '快递单号'),
  92. ('shipment_id', 'SHIPMENT ID'),
  93. ('goods_attribute', '商品属性'),
  94. ('sku', 'SKU'),
  95. ('sales_user', '商务经理'),
  96. ('service_user', '客户经理'),
  97. ('department_name', '事业部'),
  98. ('remark', '订单备注'),
  99. ('importer_name', '进口商'),
  100. ('warehouse_name', '交货仓库'),
  101. ('paid_status_name', '付款状态'),
  102. ]
  103. response['data']['columns'] = [
  104. {'key': key, 'name': name, 'check': True} for key, name in columns
  105. ]
  106. response['data']['records'] = [
  107. {key: '{0}-value'.format(key) for key, _name in columns}
  108. ]
  109. return response
  110. class McpProtocolTest(unittest.TestCase):
  111. def build_handler(self, api_client=None):
  112. token_store = InMemoryTokenStore(refresh_skew_seconds=60)
  113. token_store.save('MT_demo', '2099-01-01T00:00:00')
  114. app = GatewayApp(
  115. auth_client=None,
  116. api_client=api_client or DummyApiClient(),
  117. token_store=token_store,
  118. )
  119. return McpProtocolHandler(app)
  120. def test_initialize_returns_server_capabilities(self):
  121. handler = self.build_handler()
  122. response = handler.handle_request(
  123. {
  124. 'jsonrpc': '2.0',
  125. 'id': 1,
  126. 'method': 'initialize',
  127. 'params': {
  128. 'protocolVersion': '2025-06-18',
  129. 'capabilities': {},
  130. 'clientInfo': {'name': 'workbuddy', 'version': '1.0.0'},
  131. },
  132. }
  133. )
  134. self.assertEqual('2.0', response['jsonrpc'])
  135. self.assertEqual(1, response['id'])
  136. self.assertEqual('2025-06-18', response['result']['protocolVersion'])
  137. self.assertIn('tools', response['result']['capabilities'])
  138. self.assertEqual('fms-mcp-gateway', response['result']['serverInfo']['name'])
  139. def test_initialized_notification_does_not_emit_response(self):
  140. handler = self.build_handler()
  141. response = handler.handle_message(
  142. {
  143. 'jsonrpc': '2.0',
  144. 'method': 'notifications/initialized',
  145. }
  146. )
  147. self.assertIsNone(response)
  148. def test_tools_list_returns_registered_tools(self):
  149. handler = self.build_handler()
  150. response = handler.handle_request(
  151. {
  152. 'jsonrpc': '2.0',
  153. 'id': 2,
  154. 'method': 'tools/list',
  155. 'params': {},
  156. }
  157. )
  158. self.assertEqual('2.0', response['jsonrpc'])
  159. self.assertEqual(2, response['id'])
  160. by_name = {tool['name']: tool for tool in response['result']['tools']}
  161. self.assertIn('query_order', by_name)
  162. self.assertNotIn('bind_auth_code', by_name)
  163. self.assertIn('inputSchema', by_name['query_order'])
  164. exact_schema = by_name['query_order_exact']['inputSchema']
  165. self.assertIn(
  166. '系统订单号',
  167. exact_schema['properties']['order_number']['description'],
  168. )
  169. self.assertIsInstance(
  170. exact_schema['properties']['order_numbers']['examples'][0],
  171. list,
  172. )
  173. def test_tools_call_wraps_gateway_result_as_structured_content(self):
  174. handler = self.build_handler()
  175. response = handler.handle_request(
  176. {
  177. 'jsonrpc': '2.0',
  178. 'id': 3,
  179. 'method': 'tools/call',
  180. 'params': {
  181. 'name': 'query_order',
  182. 'arguments': {
  183. 'keyword': 'SO20260706001',
  184. 'page': 1,
  185. 'limit': 20,
  186. },
  187. },
  188. }
  189. )
  190. self.assertEqual('2.0', response['jsonrpc'])
  191. self.assertEqual(3, response['id'])
  192. self.assertFalse(response['result']['isError'])
  193. self.assertEqual('matched 1 order', response['result']['structuredContent']['summary'])
  194. self.assertEqual('SO20260706001', response['result']['structuredContent']['records'][0]['order_no'])
  195. self.assertEqual('text', response['result']['content'][0]['type'])
  196. self.assertIn('matched 1 order', response['result']['content'][0]['text'])
  197. self.assertTrue(
  198. response['result']['structuredContent']['meta']['request_id'].startswith('rq_')
  199. )
  200. def test_business_error_is_tool_error_and_stdio_continues(self):
  201. handler = self.build_handler(BusinessErrorApiClient())
  202. stdin = io.StringIO(
  203. json.dumps({
  204. 'jsonrpc': '2.0',
  205. 'id': 5,
  206. 'method': 'tools/call',
  207. 'params': {
  208. 'name': 'query_order',
  209. 'arguments': {'keyword': 'SO20260706001'},
  210. },
  211. })
  212. + '\n'
  213. + json.dumps({
  214. 'jsonrpc': '2.0',
  215. 'id': 6,
  216. 'method': 'initialize',
  217. 'params': {},
  218. })
  219. + '\n'
  220. )
  221. stdout = io.StringIO()
  222. exit_code = handler.run_stdio(stdin=stdin, stdout=stdout)
  223. responses = [json.loads(line) for line in stdout.getvalue().splitlines()]
  224. self.assertEqual(0, exit_code)
  225. self.assertEqual(2, len(responses))
  226. self.assertNotIn('error', responses[0])
  227. self.assertTrue(responses[0]['result']['isError'])
  228. self.assertIn('MCP_1301', responses[0]['result']['content'][0]['text'])
  229. self.assertEqual(
  230. 'MCP_1301',
  231. responses[0]['result']['structuredContent']['code'],
  232. )
  233. self.assertEqual(
  234. 'no order query permission',
  235. responses[0]['result']['structuredContent']['msg'],
  236. )
  237. self.assertTrue(
  238. responses[0]['result']['structuredContent']['meta']['request_id'].startswith('rq_')
  239. )
  240. self.assertEqual('2025-06-18', responses[1]['result']['protocolVersion'])
  241. def test_tools_call_renders_all_query_order_columns_in_text_content(self):
  242. token_store = InMemoryTokenStore(refresh_skew_seconds=60)
  243. token_store.save('MT_demo', '2099-01-01T00:00:00')
  244. app = GatewayApp(
  245. auth_client=None,
  246. api_client=FullColumnsApiClient(),
  247. token_store=token_store,
  248. )
  249. handler = McpProtocolHandler(app)
  250. response = handler.handle_request(
  251. {
  252. 'jsonrpc': '2.0',
  253. 'id': 4,
  254. 'method': 'tools/call',
  255. 'params': {
  256. 'name': 'query_order',
  257. 'arguments': {
  258. 'keyword': 'SO20260706001',
  259. },
  260. },
  261. }
  262. )
  263. text = response['result']['content'][0]['text']
  264. self.assertIn('表头共 40 列', text)
  265. self.assertIn('1. 订单号 (order_number)', text)
  266. self.assertIn('40. 付款状态 (paid_status_name)', text)
  267. self.assertIn('- 付款状态: paid_status_name-value', text)
  268. self.assertEqual(40, len(response['result']['structuredContent']['columns']))
  269. def test_run_stdio_serializes_non_ascii_as_ascii_json(self):
  270. class ChineseApiClient(DummyApiClient):
  271. def call_tool(self, tool_code, route_path, payload, request_id):
  272. response = super().call_tool(tool_code, route_path, payload, request_id)
  273. response['data'] = {
  274. 'summary': '共查询到 1 条轨迹记录',
  275. 'columns': [
  276. {'key': 'status', 'name': '轨迹节点'},
  277. {'key': 'location', 'name': '轨迹地点'},
  278. {'key': 'time', 'name': '时间'},
  279. {'key': 'content', 'name': '轨迹内容'},
  280. ],
  281. 'records': [
  282. {
  283. 'status': '清关放行',
  284. 'content': '启运港放行',
  285. 'location': '宁波市',
  286. 'time': '2026-07-07 10:00:00',
  287. }
  288. ],
  289. 'tips': [],
  290. }
  291. return response
  292. token_store = InMemoryTokenStore(refresh_skew_seconds=60)
  293. token_store.save('MT_demo', '2099-01-01T00:00:00')
  294. app = GatewayApp(auth_client=None, api_client=ChineseApiClient(), token_store=token_store)
  295. handler = McpProtocolHandler(app)
  296. stdin = io.StringIO(
  297. json.dumps(
  298. {
  299. 'jsonrpc': '2.0',
  300. 'id': 9,
  301. 'method': 'tools/call',
  302. 'params': {
  303. 'name': 'query_track',
  304. 'arguments': {'order_id': 6272},
  305. },
  306. }
  307. )
  308. + '\n'
  309. )
  310. stdout = io.StringIO()
  311. handler.run_stdio(stdin=stdin, stdout=stdout)
  312. line = stdout.getvalue().strip()
  313. line.encode('ascii')
  314. self.assertIn('\\u6e05\\u5173\\u653e\\u884c', line)
  315. response = json.loads(line)
  316. self.assertIn('清关放行', response['result']['content'][0]['text'])
  317. self.assertEqual(
  318. ['轨迹节点', '轨迹地点', '时间', '轨迹内容'],
  319. [
  320. header['label']
  321. for header in response['result']['structuredContent']['headers']
  322. ],
  323. )
  324. self.assertNotIn('status', json.dumps(response['result'], ensure_ascii=False))
  325. self.assertTrue(response['result']['_meta']['request_id'].startswith('rq_'))
  326. def test_run_stdio_writes_only_request_responses(self):
  327. handler = self.build_handler()
  328. stdin = io.StringIO(
  329. json.dumps(
  330. {
  331. 'jsonrpc': '2.0',
  332. 'id': 1,
  333. 'method': 'initialize',
  334. 'params': {
  335. 'protocolVersion': '2025-06-18',
  336. 'capabilities': {},
  337. 'clientInfo': {'name': 'workbuddy', 'version': '1.0.0'},
  338. },
  339. }
  340. )
  341. + '\n'
  342. + json.dumps(
  343. {
  344. 'jsonrpc': '2.0',
  345. 'method': 'notifications/initialized',
  346. }
  347. )
  348. + '\n'
  349. )
  350. stdout = io.StringIO()
  351. handler.run_stdio(stdin=stdin, stdout=stdout)
  352. lines = [line for line in stdout.getvalue().splitlines() if line.strip()]
  353. self.assertEqual(1, len(lines))
  354. response = json.loads(lines[0])
  355. self.assertEqual(1, response['id'])
  356. self.assertEqual('2025-06-18', response['result']['protocolVersion'])
  357. if __name__ == '__main__':
  358. unittest.main()