import json import unittest from unittest.mock import Mock from public_server import PublicMcpHttpHandler, create_http_handler, extract_client_ip from utils.rate_limiter import SimpleRateLimiter class FakeContext: def __init__(self, gateway_session_id): self.gateway_session_id = gateway_session_id def has_session(self): return bool(self.gateway_session_id) class FakeParser: def parse(self, headers): return FakeContext(headers.get('X-Gateway-Session', '')) class FakeGateway: def __init__(self): self.calls = [] self.list_calls = [] self.tool_result = {'code': 'MCP_0000', 'data': {'ok': True}} def registered_tool_names(self): return ('query_order', 'query_track') def list_tools(self, gateway_session_id, request_id=''): self.list_calls.append((gateway_session_id, request_id)) return [{'name': 'query_track', 'description': 'query track', 'input_schema': {'type': 'object'}}] def call_tool(self, gateway_session_id, name, arguments=None, request_id='', client_ip=''): self.calls.append((gateway_session_id, name, arguments, request_id, client_ip)) return self.tool_result class PublicMcpHttpHandlerTest(unittest.TestCase): def test_constructor_uses_registered_names_without_loading_dynamic_list(self): gateway = FakeGateway() PublicMcpHttpHandler(gateway, context_parser=FakeParser()) self.assertEqual([], gateway.list_calls) def test_handle_tools_list_passes_session_and_returns_dynamic_tools(self): gateway = FakeGateway() handler = PublicMcpHttpHandler(gateway, context_parser=FakeParser()) response = handler.handle_json_rpc( headers={'X-Gateway-Session': 'GWS_A'}, message={ 'jsonrpc': '2.0', 'id': 2, 'method': 'tools/list', 'params': {}, }, client_ip='10.0.0.5', ) self.assertEqual('GWS_A', gateway.list_calls[0][0]) self.assertTrue(gateway.list_calls[0][1].startswith('rq_http_')) self.assertEqual( ['query_track'], [tool['name'] for tool in response['result']['tools']], ) self.assertIn('inputSchema', response['result']['tools'][0]) def test_missing_session_on_tools_list_returns_protocol_error(self): gateway = FakeGateway() handler = PublicMcpHttpHandler(gateway, context_parser=FakeParser()) with self.assertLogs('public_server', level='WARNING') as logs: response = handler.handle_json_rpc( headers={}, message={ 'jsonrpc': '2.0', 'id': 3, 'method': 'tools/list', 'params': {}, }, client_ip='10.0.0.5', ) self.assertEqual(-32001, response['error']['code']) self.assertIn('Workbuddy', response['error']['message']) self.assertTrue(response['error']['data']['request_id'].startswith('rq_http_')) self.assertEqual([], gateway.list_calls) record = logs.records[0] self.assertEqual(3, record.jsonrpc_id) self.assertEqual(response['error']['data']['request_id'], record.request_id) self.assertEqual('tools/list', record.protocol_method) self.assertEqual(-32001, record.protocol_code) self.assertEqual('GATEWAY_SESSION_NOT_FOUND', record.diagnostic_reason) def test_client_request_id_is_logged_only_as_safe_hash(self): gateway = FakeGateway() handler = PublicMcpHttpHandler(gateway, context_parser=FakeParser()) value = handler._client_request_id({ 'X-Request-Id': 'client-id\nAuthorization: secret', }) self.assertRegex(value, r'^[0-9a-f]{16}$') self.assertNotIn('client-id', value) def test_tools_list_internal_exception_is_sanitized(self): class ExplodingGateway(FakeGateway): def list_tools(self, gateway_session_id, request_id=''): raise RuntimeError('database password leaked') handler = PublicMcpHttpHandler(ExplodingGateway(), context_parser=FakeParser()) response = handler.handle_json_rpc( headers={'X-Gateway-Session': 'GWS_A'}, message={'jsonrpc': '2.0', 'id': 31, 'method': 'tools/list', 'params': {}}, client_ip='10.0.0.5', ) 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_handle_tools_call_passes_gateway_session_to_public_gateway(self): gateway = FakeGateway() handler = PublicMcpHttpHandler(gateway, context_parser=FakeParser()) response = handler.handle_json_rpc( headers={ 'X-Gateway-Session': 'GWS_A', 'X-Request-Id': 'rq_public_incoming', }, message={ 'jsonrpc': '2.0', 'id': 1, 'method': 'tools/call', 'params': { 'name': 'query_order', 'arguments': {'keyword': 'USC'}, }, }, client_ip='10.0.0.5' ) self.assertEqual(False, response['result']['isError']) self.assertEqual('GWS_A', gateway.calls[0][0]) self.assertEqual('query_order', gateway.calls[0][1]) self.assertTrue(gateway.calls[0][3].startswith('rq_http_')) self.assertNotEqual('rq_public_incoming', gateway.calls[0][3]) self.assertEqual('10.0.0.5', gateway.calls[0][4]) def test_reused_client_request_id_gets_unique_server_trace_ids(self): gateway = FakeGateway() handler = PublicMcpHttpHandler(gateway, context_parser=FakeParser()) headers = { 'X-Gateway-Session': 'GWS_A', 'X-Request-Id': 'client-retry-id', } message = { 'jsonrpc': '2.0', 'id': 1, 'method': 'tools/call', 'params': {'name': 'query_track', 'arguments': {}}, } handler.handle_json_rpc(headers, message, client_ip='10.0.0.5') handler.handle_json_rpc(headers, message, client_ip='10.0.0.5') first_request_id = gateway.calls[0][3] second_request_id = gateway.calls[1][3] self.assertTrue(first_request_id.startswith('rq_http_')) self.assertTrue(second_request_id.startswith('rq_http_')) self.assertNotEqual(first_request_id, second_request_id) def test_backend_business_error_returns_tool_error_result(self): gateway = FakeGateway() gateway.tool_result = { 'code': 'MCP_1401', 'msg': 'invalid exact order query conditions', 'data': [], 'meta': {'request_id': 'rq_backend'}, } handler = PublicMcpHttpHandler(gateway, context_parser=FakeParser()) response = handler.handle_json_rpc( headers={'X-Gateway-Session': 'GWS_A'}, message={ 'jsonrpc': '2.0', 'id': 7, 'method': 'tools/call', 'params': { 'name': 'query_order', 'arguments': {'keyword': 'USC'}, }, }, client_ip='10.0.0.5', ) self.assertNotIn('error', response) self.assertTrue(response['result']['isError']) self.assertIn('MCP_1401', response['result']['content'][0]['text']) self.assertEqual( { 'code': 'MCP_1401', 'msg': 'invalid exact order query conditions', 'meta': {'request_id': 'rq_backend'}, }, response['result']['structuredContent'], ) def test_query_track_success_uses_safe_public_output(self): gateway = FakeGateway() gateway.tool_result = { 'code': 'MCP_0000', 'data': { 'summary': '共 1 条轨迹', 'columns': [ {'key': 'status', 'name': '轨迹节点'}, {'key': 'tracking_number', 'name': '快递单号'}, ], 'records': [{ 'status': '已发货', 'tracking_number': 'TN-PUBLIC', 'time_zone': '8.00', }], }, 'meta': {'request_id': 'rq_public', 'page': 1, 'limit': 5}, } handler = PublicMcpHttpHandler(gateway, context_parser=FakeParser()) response = handler.handle_json_rpc( headers={'X-Gateway-Session': 'GWS_A'}, message={ 'jsonrpc': '2.0', 'id': 8, 'method': 'tools/call', 'params': { 'name': 'query_track', 'arguments': {'tracking_number': 'TN-PUBLIC'}, }, }, client_ip='10.0.0.5', ) serialized = json.dumps(response['result'], ensure_ascii=False) self.assertFalse(response['result']['isError']) self.assertNotIn('tracking_number', serialized) self.assertNotIn('time_zone', serialized) self.assertIn('快递单号', serialized) self.assertIn('TN-PUBLIC', serialized) self.assertEqual('rq_public', response['result']['_meta']['request_id']) def test_missing_session_on_tool_call_returns_device_protocol_error(self): gateway = FakeGateway() handler = PublicMcpHttpHandler(gateway, context_parser=FakeParser()) response = handler.handle_json_rpc( headers={}, message={ 'jsonrpc': '2.0', 'id': 1, 'method': 'tools/call', 'params': {'name': 'query_order', 'arguments': {'keyword': 'USC'}}, }, ) self.assertEqual(-32001, response['error']['code']) self.assertIn('Workbuddy', response['error']['message']) self.assertEqual([], gateway.calls) def test_non_object_tool_params_fail_safely(self): gateway = FakeGateway() handler = PublicMcpHttpHandler(gateway, context_parser=FakeParser()) with self.assertLogs('public_server', level='ERROR') as logs: response = handler.handle_json_rpc( headers={'X-Gateway-Session': 'GWS_A'}, message={ 'jsonrpc': '2.0', 'id': 9, 'method': 'tools/call', 'params': 'not-an-object', }, client_ip='10.0.0.5', ) self.assertTrue(response['result']['isError']) self.assertEqual( '工具返回格式异常', response['result']['structuredContent']['message'], ) self.assertEqual([], gateway.calls) record = logs.records[0] self.assertEqual(9, record.jsonrpc_id) self.assertTrue(record.request_id.startswith('rq_http_')) self.assertEqual('', record.tool_code) self.assertEqual('MCP_9001', record.response_code) self.assertEqual('PARAM_VALIDATION_FAILED', record.diagnostic_reason) self.assertEqual('ValueError', record.exception_class) self.assertEqual(record.request_id, response['result']['_meta']['request_id']) def test_extract_client_ip_ignores_spoofable_forwarded_for_header(self): client_ip = extract_client_ip( {'X-Forwarded-For': '203.0.113.9'}, ('10.0.0.5', 54321), ) self.assertEqual('10.0.0.5', client_ip) class RateLimitTest(unittest.TestCase): def _make_handler(self, max_requests=2, max_in_flight=2): gateway = FakeGateway() limiter = SimpleRateLimiter( max_requests=max_requests, window_seconds=60, max_in_flight=max_in_flight, ) return PublicMcpHttpHandler(gateway, context_parser=FakeParser(), rate_limiter=limiter) def _tools_call_msg(self, tool_name='query_order'): return { 'jsonrpc': '2.0', 'id': 1, 'method': 'tools/call', 'params': {'name': tool_name, 'arguments': {'keyword': 'test'}}, } # --- tools/call per-tool独立限流 --- def test_known_tool_rate_limited_after_quota_exhausted(self): handler = self._make_handler(max_requests=1) handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg(), client_ip='1.2.3.4') with self.assertLogs('public_server', level='WARNING') as logs: response = handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg(), client_ip='1.2.3.4') self.assertIn('error', response) self.assertEqual(-32029, response['error']['code']) self.assertIn('Rate limit', response['error']['message']) record = logs.records[0] self.assertEqual(1, record.jsonrpc_id) self.assertEqual('query_order', record.tool_code) self.assertEqual(-32029, record.protocol_code) self.assertEqual('RATE_LIMIT_EXCEEDED', record.diagnostic_reason) self.assertEqual(response['error']['data']['request_id'], record.request_id) def test_two_known_tools_have_independent_quotas(self): # query_order 限流不影响 query_track handler = self._make_handler(max_requests=1) handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg('query_order'), client_ip='1.2.3.4') # query_order quota exhausted, query_track should still work response = handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg('query_track'), client_ip='1.2.3.4') self.assertNotIn('error', response) def test_unknown_tool_name_uses_shared_session_bucket_not_new_bucket(self): # 未知工具名应归入 session_id 桶,同一会话的不同未知工具共享同一配额,不能无限新建桶 handler = self._make_handler(max_requests=1) # 先用 session 桶打一次(用未知工具名触发) handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg('fake_tool_1'), client_ip='1.2.3.4') # 再用同会话的另一个未知工具名,应该命中同一个 session 桶,被限流 response = handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg('fake_tool_2'), client_ip='1.2.3.4') self.assertIn('error', response) self.assertIn('Rate limit', response['error']['message']) def test_different_sessions_have_independent_quotas(self): # 不同会话(不同员工)即使来自同一IP,也拥有独立的配额 handler = self._make_handler(max_requests=1) handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg(), client_ip='1.1.1.1') # 不同 session,即使同一 IP,quota 独立 → 应该允许 response = handler.handle_json_rpc({'X-Gateway-Session': 'GWS_B'}, self._tools_call_msg(), client_ip='1.1.1.1') self.assertNotIn('error', response) def test_same_session_from_different_ips_shares_quota(self): # 同一会话从不同IP发来(如移动网络切换),仍共享同一会话配额 handler = self._make_handler(max_requests=1) handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg(), client_ip='1.1.1.1') # 同 session,不同 IP → 命中同一 session 桶,被限流 response = handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg(), client_ip='2.2.2.2') self.assertIn('error', response) self.assertIn('Rate limit', response['error']['message']) # --- initialize 和 tools/list 不受限流 --- def test_initialize_is_never_rate_limited(self): # initialize 是握手协议,无论请求多少次都不应被限流 handler = self._make_handler(max_requests=1) for i in range(5): response = handler.handle_json_rpc({}, {'jsonrpc': '2.0', 'id': i, 'method': 'initialize'}, client_ip='1.2.3.4') self.assertNotIn('error', response, f'initialize should never be rate limited (attempt {i})') self.assertIn('protocolVersion', response['result']) def test_tools_list_is_never_rate_limited(self): handler = self._make_handler(max_requests=1) headers = {'X-Gateway-Session': 'GWS_A'} for request_id in range(1, 4): response = handler.handle_json_rpc( headers, {'jsonrpc': '2.0', 'id': request_id, 'method': 'tools/list'}, client_ip='1.2.3.4', ) self.assertNotIn('error', response) def test_tools_list_requests_do_not_consume_tools_call_quota(self): # tools/list 不进入限流器,因此不会消耗 tools/call 配额 handler = self._make_handler(max_requests=1) handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, {'jsonrpc': '2.0', 'id': 1, 'method': 'tools/list'}, client_ip='1.2.3.4') # tools/call 应仍可正常执行(使用 session_id:tool_name bucket) response = handler.handle_json_rpc( {'X-Gateway-Session': 'GWS_A'}, {'jsonrpc': '2.0', 'id': 2, 'method': 'tools/call', 'params': {'name': 'query_order', 'arguments': {'keyword': 'test'}}}, client_ip='1.2.3.4', ) self.assertNotIn('error', response) def test_completed_tool_call_releases_in_flight_slot_immediately(self): handler = self._make_handler(max_requests=0, max_in_flight=1) headers = {'X-Gateway-Session': 'GWS_A'} for request_id in range(1, 6): message = self._tools_call_msg('query_order') message['id'] = request_id response = handler.handle_json_rpc( headers, message, client_ip='1.2.3.4', ) self.assertNotIn('error', response) def test_exhausted_in_flight_slots_return_rate_limit_error(self): handler = self._make_handler(max_requests=0, max_in_flight=2) key = 'GWS_A:query_order' self.assertTrue(handler.rate_limiter.try_acquire(key)) self.assertTrue(handler.rate_limiter.try_acquire(key)) response = handler.handle_json_rpc( {'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg('query_order'), client_ip='1.2.3.4', ) self.assertEqual(-32029, response['error']['code']) self.assertIn('in progress', response['error']['message']) def test_failed_tool_call_releases_in_flight_slot(self): handler = self._make_handler(max_requests=0, max_in_flight=1) original_call = handler.gateway_app.call_tool def fail(*args, **kwargs): raise RuntimeError('backend failed') handler.gateway_app.call_tool = fail handler.handle_json_rpc( {'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg('query_order'), client_ip='1.2.3.4', ) handler.gateway_app.call_tool = original_call response = handler.handle_json_rpc( {'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg('query_order'), client_ip='1.2.3.4', ) self.assertNotIn('error', response) class HttpDisconnectTest(unittest.TestCase): def test_broken_pipe_while_writing_response_is_handled(self): handler_class = create_http_handler(FakeGateway(), rate_limiter=None) handler = object.__new__(handler_class) handler.send_response = Mock() handler.send_header = Mock() handler.end_headers = Mock() handler.wfile = Mock() handler.wfile.write.side_effect = BrokenPipeError() with self.assertLogs('public_server', level='INFO') as logs: handler._write_json({'jsonrpc': '2.0', 'id': 1, 'result': {}}) self.assertIn('client disconnected before response', logs.output[0]) if __name__ == '__main__': unittest.main()