test_mcp_protocol.py 16 KB

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