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()