test_mcp_protocol.py 19 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533
  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, request_id=''):
  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, reporter=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, reporter=reporter)
  120. def test_stdio_tool_call_emits_correlated_diagnostic_stages(self):
  121. reporter = RecordingReporter()
  122. handler = self.build_handler(reporter=reporter)
  123. response = handler.handle_request({
  124. 'jsonrpc': '2.0',
  125. 'id': 10,
  126. 'method': 'tools/call',
  127. 'params': {
  128. 'name': 'query_order',
  129. 'arguments': {'keyword': 'SO20260706001'},
  130. },
  131. })
  132. self.assertFalse(response['result']['isError'])
  133. self.assertEqual(
  134. [
  135. 'request_ingress',
  136. 'protocol_validation',
  137. 'backend_call',
  138. 'backend_call',
  139. 'response_safety',
  140. ],
  141. [event['stage'] for event in reporter.events],
  142. )
  143. self.assertEqual(1, len({event['request_id'] for event in reporter.events}))
  144. def test_stdio_missing_tool_name_returns_invalid_params(self):
  145. reporter = RecordingReporter()
  146. response = self.build_handler(reporter=reporter).handle_request({
  147. 'jsonrpc': '2.0',
  148. 'id': 11,
  149. 'method': 'tools/call',
  150. 'params': {'arguments': {}},
  151. })
  152. self.assertIn('error', response)
  153. self.assertEqual(-32602, response['error']['code'])
  154. self.assertNotIn('result', response)
  155. self.assertEqual(
  156. ['request_ingress', 'protocol_validation'],
  157. [event['stage'] for event in reporter.events],
  158. )
  159. self.assertEqual('failed', reporter.events[-1]['status'])
  160. self.assertEqual(
  161. 'PARAM_VALIDATION_FAILED',
  162. reporter.events[-1]['event_code'],
  163. )
  164. def test_stdio_reporter_failure_does_not_change_response(self):
  165. class BrokenReporter:
  166. def report(self, _event):
  167. raise RuntimeError('support unavailable')
  168. response = self.build_handler(reporter=BrokenReporter()).handle_request({
  169. 'jsonrpc': '2.0',
  170. 'id': 1,
  171. 'method': 'initialize',
  172. 'params': {},
  173. })
  174. self.assertIn('result', response)
  175. def test_stdio_invalid_tool_result_emits_response_safety_failure(self):
  176. class InvalidGateway:
  177. def call_tool(self, name, arguments, request_id=''):
  178. return []
  179. reporter = RecordingReporter()
  180. response = McpProtocolHandler(
  181. InvalidGateway(),
  182. reporter=reporter,
  183. ).handle_request({
  184. 'jsonrpc': '2.0',
  185. 'id': 1,
  186. 'method': 'tools/call',
  187. 'params': {'name': 'query_order', 'arguments': {}},
  188. })
  189. self.assertTrue(response['result']['isError'])
  190. event = next(
  191. item for item in reporter.events
  192. if item['stage'] == 'response_safety'
  193. )
  194. self.assertEqual('failed', event['status'])
  195. def test_non_tool_exception_is_sanitized(self):
  196. class ExplodingGateway:
  197. def list_tools(self):
  198. raise RuntimeError('database password leaked')
  199. response = McpProtocolHandler(ExplodingGateway()).handle_request({
  200. 'jsonrpc': '2.0',
  201. 'id': 99,
  202. 'method': 'tools/list',
  203. })
  204. self.assertEqual(-32000, response['error']['code'])
  205. self.assertEqual('Gateway request failed. Please try again later.', response['error']['message'])
  206. self.assertNotIn('password', json.dumps(response))
  207. def test_tool_exception_is_logged_and_correlated_without_raw_message(self):
  208. class ExplodingGateway:
  209. def call_tool(self, name, arguments, request_id=''):
  210. raise RuntimeError('database password leaked')
  211. handler = McpProtocolHandler(ExplodingGateway())
  212. with self.assertLogs('mcp_protocol', level='ERROR') as logs:
  213. response = handler.handle_request({
  214. 'jsonrpc': '2.0',
  215. 'id': 100,
  216. 'method': 'tools/call',
  217. 'params': {'name': 'query_track', 'arguments': {'tracking_number': 'TN1'}},
  218. })
  219. serialized = json.dumps(response)
  220. self.assertTrue(response['result']['isError'])
  221. self.assertNotIn('password', serialized)
  222. record = logs.records[0]
  223. self.assertEqual(100, record.jsonrpc_id)
  224. self.assertTrue(record.request_id.startswith('rq_stdio_'))
  225. self.assertEqual('query_track', record.tool_code)
  226. self.assertEqual('MCP_9001', record.response_code)
  227. self.assertEqual('UNEXPECTED_EXCEPTION', record.diagnostic_reason)
  228. self.assertEqual('RuntimeError', record.exception_class)
  229. self.assertEqual(record.request_id, response['result']['_meta']['request_id'])
  230. def test_initialize_returns_server_capabilities(self):
  231. handler = self.build_handler()
  232. response = handler.handle_request(
  233. {
  234. 'jsonrpc': '2.0',
  235. 'id': 1,
  236. 'method': 'initialize',
  237. 'params': {
  238. 'protocolVersion': '2025-06-18',
  239. 'capabilities': {},
  240. 'clientInfo': {'name': 'workbuddy', 'version': '1.0.0'},
  241. },
  242. }
  243. )
  244. self.assertEqual('2.0', response['jsonrpc'])
  245. self.assertEqual(1, response['id'])
  246. self.assertEqual('2025-06-18', response['result']['protocolVersion'])
  247. self.assertIn('tools', response['result']['capabilities'])
  248. self.assertEqual('fms-mcp-gateway', response['result']['serverInfo']['name'])
  249. def test_initialized_notification_does_not_emit_response(self):
  250. handler = self.build_handler()
  251. response = handler.handle_message(
  252. {
  253. 'jsonrpc': '2.0',
  254. 'method': 'notifications/initialized',
  255. }
  256. )
  257. self.assertIsNone(response)
  258. def test_tools_list_returns_registered_tools(self):
  259. handler = self.build_handler()
  260. response = handler.handle_request(
  261. {
  262. 'jsonrpc': '2.0',
  263. 'id': 2,
  264. 'method': 'tools/list',
  265. 'params': {},
  266. }
  267. )
  268. self.assertEqual('2.0', response['jsonrpc'])
  269. self.assertEqual(2, response['id'])
  270. by_name = {tool['name']: tool for tool in response['result']['tools']}
  271. self.assertIn('query_order', by_name)
  272. self.assertNotIn('bind_auth_code', by_name)
  273. self.assertIn('inputSchema', by_name['query_order'])
  274. exact_schema = by_name['query_order_exact']['inputSchema']
  275. self.assertIn(
  276. '系统订单号',
  277. exact_schema['properties']['order_number']['description'],
  278. )
  279. self.assertIsInstance(
  280. exact_schema['properties']['order_numbers']['examples'][0],
  281. list,
  282. )
  283. def test_tools_call_wraps_gateway_result_as_structured_content(self):
  284. handler = self.build_handler()
  285. response = handler.handle_request(
  286. {
  287. 'jsonrpc': '2.0',
  288. 'id': 3,
  289. 'method': 'tools/call',
  290. 'params': {
  291. 'name': 'query_order',
  292. 'arguments': {
  293. 'keyword': 'SO20260706001',
  294. 'page': 1,
  295. 'limit': 20,
  296. },
  297. },
  298. }
  299. )
  300. self.assertEqual('2.0', response['jsonrpc'])
  301. self.assertEqual(3, response['id'])
  302. self.assertFalse(response['result']['isError'])
  303. self.assertEqual('matched 1 order', response['result']['structuredContent']['summary'])
  304. self.assertEqual('SO20260706001', response['result']['structuredContent']['records'][0]['order_no'])
  305. self.assertEqual('text', response['result']['content'][0]['type'])
  306. self.assertIn('matched 1 order', response['result']['content'][0]['text'])
  307. self.assertTrue(
  308. response['result']['structuredContent']['meta']['request_id'].startswith('rq_')
  309. )
  310. def test_business_error_is_tool_error_and_stdio_continues(self):
  311. handler = self.build_handler(BusinessErrorApiClient())
  312. stdin = io.StringIO(
  313. json.dumps({
  314. 'jsonrpc': '2.0',
  315. 'id': 5,
  316. 'method': 'tools/call',
  317. 'params': {
  318. 'name': 'query_order',
  319. 'arguments': {'keyword': 'SO20260706001'},
  320. },
  321. })
  322. + '\n'
  323. + json.dumps({
  324. 'jsonrpc': '2.0',
  325. 'id': 6,
  326. 'method': 'initialize',
  327. 'params': {},
  328. })
  329. + '\n'
  330. )
  331. stdout = io.StringIO()
  332. exit_code = handler.run_stdio(stdin=stdin, stdout=stdout)
  333. responses = [json.loads(line) for line in stdout.getvalue().splitlines()]
  334. self.assertEqual(0, exit_code)
  335. self.assertEqual(2, len(responses))
  336. self.assertNotIn('error', responses[0])
  337. self.assertTrue(responses[0]['result']['isError'])
  338. self.assertIn('MCP_1301', responses[0]['result']['content'][0]['text'])
  339. self.assertEqual(
  340. 'MCP_1301',
  341. responses[0]['result']['structuredContent']['code'],
  342. )
  343. self.assertEqual(
  344. 'no order query permission',
  345. responses[0]['result']['structuredContent']['msg'],
  346. )
  347. self.assertTrue(
  348. responses[0]['result']['structuredContent']['meta']['request_id'].startswith('rq_')
  349. )
  350. self.assertEqual('2025-06-18', responses[1]['result']['protocolVersion'])
  351. def test_tools_call_renders_all_query_order_columns_in_text_content(self):
  352. token_store = InMemoryTokenStore(refresh_skew_seconds=60)
  353. token_store.save('MT_demo', '2099-01-01T00:00:00')
  354. app = GatewayApp(
  355. auth_client=None,
  356. api_client=FullColumnsApiClient(),
  357. token_store=token_store,
  358. )
  359. handler = McpProtocolHandler(app)
  360. response = handler.handle_request(
  361. {
  362. 'jsonrpc': '2.0',
  363. 'id': 4,
  364. 'method': 'tools/call',
  365. 'params': {
  366. 'name': 'query_order',
  367. 'arguments': {
  368. 'keyword': 'SO20260706001',
  369. },
  370. },
  371. }
  372. )
  373. text = response['result']['content'][0]['text']
  374. self.assertIn('表头共 40 列', text)
  375. self.assertIn('1. 订单号 (order_number)', text)
  376. self.assertIn('40. 付款状态 (paid_status_name)', text)
  377. self.assertIn('- 付款状态: paid_status_name-value', text)
  378. self.assertEqual(40, len(response['result']['structuredContent']['columns']))
  379. def test_run_stdio_serializes_non_ascii_as_ascii_json(self):
  380. class ChineseApiClient(DummyApiClient):
  381. def call_tool(self, tool_code, route_path, payload, request_id):
  382. response = super().call_tool(tool_code, route_path, payload, request_id)
  383. response['data'] = {
  384. 'summary': '共查询到 1 条轨迹记录',
  385. 'columns': [
  386. {'key': 'status', 'name': '轨迹节点'},
  387. {'key': 'location', 'name': '轨迹地点'},
  388. {'key': 'time', 'name': '时间'},
  389. {'key': 'content', 'name': '轨迹内容'},
  390. ],
  391. 'records': [
  392. {
  393. 'status': '清关放行',
  394. 'content': '启运港放行',
  395. 'location': '宁波市',
  396. 'time': '2026-07-07 10:00:00',
  397. }
  398. ],
  399. 'tips': [],
  400. }
  401. return response
  402. token_store = InMemoryTokenStore(refresh_skew_seconds=60)
  403. token_store.save('MT_demo', '2099-01-01T00:00:00')
  404. app = GatewayApp(auth_client=None, api_client=ChineseApiClient(), token_store=token_store)
  405. handler = McpProtocolHandler(app)
  406. stdin = io.StringIO(
  407. json.dumps(
  408. {
  409. 'jsonrpc': '2.0',
  410. 'id': 9,
  411. 'method': 'tools/call',
  412. 'params': {
  413. 'name': 'query_track',
  414. 'arguments': {'order_id': 6272},
  415. },
  416. }
  417. )
  418. + '\n'
  419. )
  420. stdout = io.StringIO()
  421. handler.run_stdio(stdin=stdin, stdout=stdout)
  422. line = stdout.getvalue().strip()
  423. line.encode('ascii')
  424. self.assertIn('\\u6e05\\u5173\\u653e\\u884c', line)
  425. response = json.loads(line)
  426. self.assertIn('清关放行', response['result']['content'][0]['text'])
  427. self.assertEqual(
  428. ['轨迹节点', '轨迹地点', '时间', '轨迹内容'],
  429. [
  430. header['label']
  431. for header in response['result']['structuredContent']['headers']
  432. ],
  433. )
  434. self.assertNotIn('status', json.dumps(response['result'], ensure_ascii=False))
  435. self.assertTrue(response['result']['_meta']['request_id'].startswith('rq_'))
  436. def test_run_stdio_writes_only_request_responses(self):
  437. handler = self.build_handler()
  438. stdin = io.StringIO(
  439. json.dumps(
  440. {
  441. 'jsonrpc': '2.0',
  442. 'id': 1,
  443. 'method': 'initialize',
  444. 'params': {
  445. 'protocolVersion': '2025-06-18',
  446. 'capabilities': {},
  447. 'clientInfo': {'name': 'workbuddy', 'version': '1.0.0'},
  448. },
  449. }
  450. )
  451. + '\n'
  452. + json.dumps(
  453. {
  454. 'jsonrpc': '2.0',
  455. 'method': 'notifications/initialized',
  456. }
  457. )
  458. + '\n'
  459. )
  460. stdout = io.StringIO()
  461. handler.run_stdio(stdin=stdin, stdout=stdout)
  462. lines = [line for line in stdout.getvalue().splitlines() if line.strip()]
  463. self.assertEqual(1, len(lines))
  464. response = json.loads(lines[0])
  465. self.assertEqual(1, response['id'])
  466. self.assertEqual('2025-06-18', response['result']['protocolVersion'])
  467. class RecordingReporter:
  468. def __init__(self):
  469. self.events = []
  470. def report(self, event):
  471. self.events.append(event)
  472. return True
  473. if __name__ == '__main__':
  474. unittest.main()