| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533 |
- import io
- import json
- import unittest
- from app import GatewayApp
- from services.token_store import InMemoryTokenStore
- from mcp_protocol import McpProtocolHandler
- class DummyApiClient:
- def __init__(self):
- self.calls = []
- def list_enabled_tools(self, request_id=''):
- return {
- 'code': 'MCP_0000',
- 'data': {
- 'tool_codes': [
- 'query_order',
- 'query_track',
- 'query_order_exact',
- 'list_order_filter_options',
- ],
- },
- }
- def call_tool(self, tool_code, route_path, payload, request_id):
- self.calls.append(
- {
- 'tool_code': tool_code,
- 'route_path': route_path,
- 'payload': payload,
- 'request_id': request_id,
- }
- )
- return {
- 'code': 'MCP_0000',
- 'msg': 'success',
- 'data': {
- 'summary': 'matched 1 order',
- 'records': [
- {
- 'order_no': 'SO20260706001',
- }
- ],
- 'tips': ['scoped by employee permissions'],
- },
- 'meta': {
- 'request_id': request_id,
- },
- }
- class BusinessErrorApiClient(DummyApiClient):
- def call_tool(self, tool_code, route_path, payload, request_id):
- super().call_tool(tool_code, route_path, payload, request_id)
- return {
- 'code': 'MCP_1301',
- 'msg': 'no order query permission',
- 'data': [],
- 'meta': {
- 'request_id': request_id,
- },
- }
- class FullColumnsApiClient(DummyApiClient):
- def call_tool(self, tool_code, route_path, payload, request_id):
- response = super().call_tool(tool_code, route_path, payload, request_id)
- columns = [
- ('order_number', '订单号'),
- ('reference_number', '客户参考号'),
- ('status_txt_name', '状态'),
- ('check_status_txt_name', '是否已查验'),
- ('customer_name', '客户名称'),
- ('customer_account_type_name', '客户属性'),
- ('inbound_date', '入库时间'),
- ('wo_num', '未完成工单'),
- ('product_name', '物流产品'),
- ('inbound_pieces', '件数'),
- ('inbound_volume', '体积(CBM)'),
- ('inbound_weight', '重量(KG)'),
- ('pro_cn_name', '品名'),
- ('export_declaration_type', '报关方式'),
- ('merge_declare_number', '合并报关单号'),
- ('delivery_address', '派送地址'),
- ('container_code', '柜号'),
- ('out_status_txt', '排舱单状态'),
- ('hinge_of_destination', '目的港'),
- ('etd', 'ETD'),
- ('atd', 'ATD'),
- ('eta', 'ETA'),
- ('ata', 'ATA'),
- ('release_time', '清关放行时间'),
- ('oversea_inbound_date', '海外入库时间'),
- ('appt_time', 'APPT时间'),
- ('est_loading_time', '预计装柜时间'),
- ('pickup_time', '海外提柜时间'),
- ('delivery_way_title', '派送方式'),
- ('tracking_number', '快递单号'),
- ('shipment_id', 'SHIPMENT ID'),
- ('goods_attribute', '商品属性'),
- ('sku', 'SKU'),
- ('sales_user', '商务经理'),
- ('service_user', '客户经理'),
- ('department_name', '事业部'),
- ('remark', '订单备注'),
- ('importer_name', '进口商'),
- ('warehouse_name', '交货仓库'),
- ('paid_status_name', '付款状态'),
- ]
- response['data']['columns'] = [
- {'key': key, 'name': name, 'check': True} for key, name in columns
- ]
- response['data']['records'] = [
- {key: '{0}-value'.format(key) for key, _name in columns}
- ]
- return response
- class McpProtocolTest(unittest.TestCase):
- def build_handler(self, api_client=None, reporter=None):
- token_store = InMemoryTokenStore(refresh_skew_seconds=60)
- token_store.save('MT_demo', '2099-01-01T00:00:00')
- app = GatewayApp(
- auth_client=None,
- api_client=api_client or DummyApiClient(),
- token_store=token_store,
- )
- return McpProtocolHandler(app, reporter=reporter)
- def test_stdio_tool_call_emits_correlated_diagnostic_stages(self):
- reporter = RecordingReporter()
- handler = self.build_handler(reporter=reporter)
- response = handler.handle_request({
- 'jsonrpc': '2.0',
- 'id': 10,
- 'method': 'tools/call',
- 'params': {
- 'name': 'query_order',
- 'arguments': {'keyword': 'SO20260706001'},
- },
- })
- self.assertFalse(response['result']['isError'])
- self.assertEqual(
- [
- 'request_ingress',
- 'protocol_validation',
- 'backend_call',
- 'backend_call',
- 'response_safety',
- ],
- [event['stage'] for event in reporter.events],
- )
- self.assertEqual(1, len({event['request_id'] for event in reporter.events}))
- def test_stdio_missing_tool_name_returns_invalid_params(self):
- reporter = RecordingReporter()
- response = self.build_handler(reporter=reporter).handle_request({
- 'jsonrpc': '2.0',
- 'id': 11,
- 'method': 'tools/call',
- 'params': {'arguments': {}},
- })
- self.assertIn('error', response)
- self.assertEqual(-32602, response['error']['code'])
- self.assertNotIn('result', response)
- self.assertEqual(
- ['request_ingress', 'protocol_validation'],
- [event['stage'] for event in reporter.events],
- )
- self.assertEqual('failed', reporter.events[-1]['status'])
- self.assertEqual(
- 'PARAM_VALIDATION_FAILED',
- reporter.events[-1]['event_code'],
- )
- def test_stdio_reporter_failure_does_not_change_response(self):
- class BrokenReporter:
- def report(self, _event):
- raise RuntimeError('support unavailable')
- response = self.build_handler(reporter=BrokenReporter()).handle_request({
- 'jsonrpc': '2.0',
- 'id': 1,
- 'method': 'initialize',
- 'params': {},
- })
- self.assertIn('result', response)
- def test_stdio_invalid_tool_result_emits_response_safety_failure(self):
- class InvalidGateway:
- def call_tool(self, name, arguments, request_id=''):
- return []
- reporter = RecordingReporter()
- response = McpProtocolHandler(
- InvalidGateway(),
- reporter=reporter,
- ).handle_request({
- 'jsonrpc': '2.0',
- 'id': 1,
- 'method': 'tools/call',
- 'params': {'name': 'query_order', 'arguments': {}},
- })
- self.assertTrue(response['result']['isError'])
- event = next(
- item for item in reporter.events
- if item['stage'] == 'response_safety'
- )
- self.assertEqual('failed', event['status'])
- def test_non_tool_exception_is_sanitized(self):
- class ExplodingGateway:
- def list_tools(self):
- raise RuntimeError('database password leaked')
- response = McpProtocolHandler(ExplodingGateway()).handle_request({
- 'jsonrpc': '2.0',
- 'id': 99,
- 'method': 'tools/list',
- })
- self.assertEqual(-32000, response['error']['code'])
- self.assertEqual('Gateway request failed. Please try again later.', response['error']['message'])
- self.assertNotIn('password', json.dumps(response))
- def test_tool_exception_is_logged_and_correlated_without_raw_message(self):
- class ExplodingGateway:
- def call_tool(self, name, arguments, request_id=''):
- raise RuntimeError('database password leaked')
- handler = McpProtocolHandler(ExplodingGateway())
- with self.assertLogs('mcp_protocol', level='ERROR') as logs:
- response = handler.handle_request({
- 'jsonrpc': '2.0',
- 'id': 100,
- 'method': 'tools/call',
- 'params': {'name': 'query_track', 'arguments': {'tracking_number': 'TN1'}},
- })
- serialized = json.dumps(response)
- self.assertTrue(response['result']['isError'])
- self.assertNotIn('password', serialized)
- record = logs.records[0]
- self.assertEqual(100, record.jsonrpc_id)
- self.assertTrue(record.request_id.startswith('rq_stdio_'))
- self.assertEqual('query_track', record.tool_code)
- self.assertEqual('MCP_9001', record.response_code)
- self.assertEqual('UNEXPECTED_EXCEPTION', record.diagnostic_reason)
- self.assertEqual('RuntimeError', record.exception_class)
- self.assertEqual(record.request_id, response['result']['_meta']['request_id'])
- def test_initialize_returns_server_capabilities(self):
- handler = self.build_handler()
- response = handler.handle_request(
- {
- 'jsonrpc': '2.0',
- 'id': 1,
- 'method': 'initialize',
- 'params': {
- 'protocolVersion': '2025-06-18',
- 'capabilities': {},
- 'clientInfo': {'name': 'workbuddy', 'version': '1.0.0'},
- },
- }
- )
- self.assertEqual('2.0', response['jsonrpc'])
- self.assertEqual(1, response['id'])
- self.assertEqual('2025-06-18', response['result']['protocolVersion'])
- self.assertIn('tools', response['result']['capabilities'])
- self.assertEqual('fms-mcp-gateway', response['result']['serverInfo']['name'])
- def test_initialized_notification_does_not_emit_response(self):
- handler = self.build_handler()
- response = handler.handle_message(
- {
- 'jsonrpc': '2.0',
- 'method': 'notifications/initialized',
- }
- )
- self.assertIsNone(response)
- def test_tools_list_returns_registered_tools(self):
- handler = self.build_handler()
- response = handler.handle_request(
- {
- 'jsonrpc': '2.0',
- 'id': 2,
- 'method': 'tools/list',
- 'params': {},
- }
- )
- self.assertEqual('2.0', response['jsonrpc'])
- self.assertEqual(2, response['id'])
- by_name = {tool['name']: tool for tool in response['result']['tools']}
- self.assertIn('query_order', by_name)
- self.assertNotIn('bind_auth_code', by_name)
- self.assertIn('inputSchema', by_name['query_order'])
- exact_schema = by_name['query_order_exact']['inputSchema']
- self.assertIn(
- '系统订单号',
- exact_schema['properties']['order_number']['description'],
- )
- self.assertIsInstance(
- exact_schema['properties']['order_numbers']['examples'][0],
- list,
- )
- def test_tools_call_wraps_gateway_result_as_structured_content(self):
- handler = self.build_handler()
- response = handler.handle_request(
- {
- 'jsonrpc': '2.0',
- 'id': 3,
- 'method': 'tools/call',
- 'params': {
- 'name': 'query_order',
- 'arguments': {
- 'keyword': 'SO20260706001',
- 'page': 1,
- 'limit': 20,
- },
- },
- }
- )
- self.assertEqual('2.0', response['jsonrpc'])
- self.assertEqual(3, response['id'])
- self.assertFalse(response['result']['isError'])
- self.assertEqual('matched 1 order', response['result']['structuredContent']['summary'])
- self.assertEqual('SO20260706001', response['result']['structuredContent']['records'][0]['order_no'])
- self.assertEqual('text', response['result']['content'][0]['type'])
- self.assertIn('matched 1 order', response['result']['content'][0]['text'])
- self.assertTrue(
- response['result']['structuredContent']['meta']['request_id'].startswith('rq_')
- )
- def test_business_error_is_tool_error_and_stdio_continues(self):
- handler = self.build_handler(BusinessErrorApiClient())
- stdin = io.StringIO(
- json.dumps({
- 'jsonrpc': '2.0',
- 'id': 5,
- 'method': 'tools/call',
- 'params': {
- 'name': 'query_order',
- 'arguments': {'keyword': 'SO20260706001'},
- },
- })
- + '\n'
- + json.dumps({
- 'jsonrpc': '2.0',
- 'id': 6,
- 'method': 'initialize',
- 'params': {},
- })
- + '\n'
- )
- stdout = io.StringIO()
- exit_code = handler.run_stdio(stdin=stdin, stdout=stdout)
- responses = [json.loads(line) for line in stdout.getvalue().splitlines()]
- self.assertEqual(0, exit_code)
- self.assertEqual(2, len(responses))
- self.assertNotIn('error', responses[0])
- self.assertTrue(responses[0]['result']['isError'])
- self.assertIn('MCP_1301', responses[0]['result']['content'][0]['text'])
- self.assertEqual(
- 'MCP_1301',
- responses[0]['result']['structuredContent']['code'],
- )
- self.assertEqual(
- 'no order query permission',
- responses[0]['result']['structuredContent']['msg'],
- )
- self.assertTrue(
- responses[0]['result']['structuredContent']['meta']['request_id'].startswith('rq_')
- )
- self.assertEqual('2025-06-18', responses[1]['result']['protocolVersion'])
- def test_tools_call_renders_all_query_order_columns_in_text_content(self):
- token_store = InMemoryTokenStore(refresh_skew_seconds=60)
- token_store.save('MT_demo', '2099-01-01T00:00:00')
- app = GatewayApp(
- auth_client=None,
- api_client=FullColumnsApiClient(),
- token_store=token_store,
- )
- handler = McpProtocolHandler(app)
- response = handler.handle_request(
- {
- 'jsonrpc': '2.0',
- 'id': 4,
- 'method': 'tools/call',
- 'params': {
- 'name': 'query_order',
- 'arguments': {
- 'keyword': 'SO20260706001',
- },
- },
- }
- )
- text = response['result']['content'][0]['text']
- self.assertIn('表头共 40 列', text)
- self.assertIn('1. 订单号 (order_number)', text)
- self.assertIn('40. 付款状态 (paid_status_name)', text)
- self.assertIn('- 付款状态: paid_status_name-value', text)
- self.assertEqual(40, len(response['result']['structuredContent']['columns']))
- def test_run_stdio_serializes_non_ascii_as_ascii_json(self):
- class ChineseApiClient(DummyApiClient):
- def call_tool(self, tool_code, route_path, payload, request_id):
- response = super().call_tool(tool_code, route_path, payload, request_id)
- response['data'] = {
- 'summary': '共查询到 1 条轨迹记录',
- 'columns': [
- {'key': 'status', 'name': '轨迹节点'},
- {'key': 'location', 'name': '轨迹地点'},
- {'key': 'time', 'name': '时间'},
- {'key': 'content', 'name': '轨迹内容'},
- ],
- 'records': [
- {
- 'status': '清关放行',
- 'content': '启运港放行',
- 'location': '宁波市',
- 'time': '2026-07-07 10:00:00',
- }
- ],
- 'tips': [],
- }
- return response
- token_store = InMemoryTokenStore(refresh_skew_seconds=60)
- token_store.save('MT_demo', '2099-01-01T00:00:00')
- app = GatewayApp(auth_client=None, api_client=ChineseApiClient(), token_store=token_store)
- handler = McpProtocolHandler(app)
- stdin = io.StringIO(
- json.dumps(
- {
- 'jsonrpc': '2.0',
- 'id': 9,
- 'method': 'tools/call',
- 'params': {
- 'name': 'query_track',
- 'arguments': {'order_id': 6272},
- },
- }
- )
- + '\n'
- )
- stdout = io.StringIO()
- handler.run_stdio(stdin=stdin, stdout=stdout)
- line = stdout.getvalue().strip()
- line.encode('ascii')
- self.assertIn('\\u6e05\\u5173\\u653e\\u884c', line)
- response = json.loads(line)
- self.assertIn('清关放行', response['result']['content'][0]['text'])
- self.assertEqual(
- ['轨迹节点', '轨迹地点', '时间', '轨迹内容'],
- [
- header['label']
- for header in response['result']['structuredContent']['headers']
- ],
- )
- self.assertNotIn('status', json.dumps(response['result'], ensure_ascii=False))
- self.assertTrue(response['result']['_meta']['request_id'].startswith('rq_'))
- def test_run_stdio_writes_only_request_responses(self):
- handler = self.build_handler()
- stdin = io.StringIO(
- json.dumps(
- {
- 'jsonrpc': '2.0',
- 'id': 1,
- 'method': 'initialize',
- 'params': {
- 'protocolVersion': '2025-06-18',
- 'capabilities': {},
- 'clientInfo': {'name': 'workbuddy', 'version': '1.0.0'},
- },
- }
- )
- + '\n'
- + json.dumps(
- {
- 'jsonrpc': '2.0',
- 'method': 'notifications/initialized',
- }
- )
- + '\n'
- )
- stdout = io.StringIO()
- handler.run_stdio(stdin=stdin, stdout=stdout)
- lines = [line for line in stdout.getvalue().splitlines() if line.strip()]
- self.assertEqual(1, len(lines))
- response = json.loads(lines[0])
- self.assertEqual(1, response['id'])
- self.assertEqual('2025-06-18', response['result']['protocolVersion'])
- class RecordingReporter:
- def __init__(self):
- self.events = []
- def report(self, event):
- self.events.append(event)
- return True
- if __name__ == '__main__':
- unittest.main()
|