|
|
@@ -1,6 +1,10 @@
|
|
|
+import json
|
|
|
import unittest
|
|
|
+from unittest.mock import Mock
|
|
|
|
|
|
-from public_server import PublicMcpHttpHandler, extract_client_ip
|
|
|
+from public_server import PublicMcpHttpHandler, create_http_handler, extract_client_ip
|
|
|
+from services.diagnostic_event import RequestDiagnosticEmitter
|
|
|
+from utils.rate_limiter import SimpleRateLimiter
|
|
|
|
|
|
|
|
|
class FakeContext:
|
|
|
@@ -19,21 +23,238 @@ class FakeParser:
|
|
|
class FakeGateway:
|
|
|
def __init__(self):
|
|
|
self.calls = []
|
|
|
+ self.list_calls = []
|
|
|
+ self.tool_result = {'code': 'MCP_0000', 'data': {'ok': True}}
|
|
|
|
|
|
- def list_tools(self):
|
|
|
- return [{'name': 'query_order', 'description': 'query order', 'input_schema': {'type': 'object'}}]
|
|
|
+ def registered_tool_names(self):
|
|
|
+ return ('query_order', 'query_track')
|
|
|
|
|
|
- def call_tool(self, gateway_session_id, name, arguments=None, request_id=''):
|
|
|
- self.calls.append((gateway_session_id, name, arguments, request_id))
|
|
|
- return {'code': 'MCP_0000', 'data': {'ok': True}}
|
|
|
+ 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='',
|
|
|
+ diagnostic_emitter=None,
|
|
|
+ ):
|
|
|
+ self.calls.append((gateway_session_id, name, arguments, request_id, client_ip))
|
|
|
+ return self.tool_result
|
|
|
|
|
|
|
|
|
class PublicMcpHttpHandlerTest(unittest.TestCase):
|
|
|
- def test_handle_tools_call_passes_gateway_session_to_public_gateway(self):
|
|
|
+ def test_initialize_emits_successful_protocol_validation(self):
|
|
|
+ reporter = RecordingReporter()
|
|
|
+ handler = PublicMcpHttpHandler(FakeGateway(), reporter=reporter)
|
|
|
+
|
|
|
+ response = handler.handle_json_rpc(
|
|
|
+ headers={},
|
|
|
+ message={'jsonrpc': '2.0', 'id': 1, 'method': 'initialize'},
|
|
|
+ )
|
|
|
+
|
|
|
+ self.assertIn('result', response)
|
|
|
+ self.assertEqual(
|
|
|
+ ['request_ingress', 'protocol_validation'],
|
|
|
+ [event['stage'] for event in reporter.events],
|
|
|
+ )
|
|
|
+ self.assertEqual('succeeded', reporter.events[-1]['status'])
|
|
|
+
|
|
|
+ def test_missing_session_emits_safe_ingress_and_session_failure(self):
|
|
|
+ reporter = RecordingReporter()
|
|
|
+ handler = PublicMcpHttpHandler(
|
|
|
+ FakeGateway(),
|
|
|
+ context_parser=FakeParser(),
|
|
|
+ reporter=reporter,
|
|
|
+ )
|
|
|
+
|
|
|
+ response = handler.handle_json_rpc(
|
|
|
+ headers={},
|
|
|
+ message={
|
|
|
+ 'jsonrpc': '2.0',
|
|
|
+ 'id': 1,
|
|
|
+ 'method': 'tools/call',
|
|
|
+ 'params': {'name': 'query_order', 'arguments': {}},
|
|
|
+ },
|
|
|
+ )
|
|
|
+
|
|
|
+ self.assertEqual(-32001, response['error']['code'])
|
|
|
+ self.assertEqual(
|
|
|
+ ['request_ingress', 'protocol_validation', 'gateway_session'],
|
|
|
+ [event['stage'] for event in reporter.events],
|
|
|
+ )
|
|
|
+ self.assertEqual('failed', reporter.events[-1]['status'])
|
|
|
+
|
|
|
+ def test_reporter_failure_does_not_change_successful_response(self):
|
|
|
+ class BrokenReporter:
|
|
|
+ def report(self, _event):
|
|
|
+ raise RuntimeError('support unavailable')
|
|
|
+
|
|
|
+ handler = PublicMcpHttpHandler(
|
|
|
+ FakeGateway(),
|
|
|
+ context_parser=FakeParser(),
|
|
|
+ reporter=BrokenReporter(),
|
|
|
+ )
|
|
|
+
|
|
|
+ response = handler.handle_json_rpc(
|
|
|
+ headers={'X-Gateway-Session': 'GWS_A'},
|
|
|
+ message={'jsonrpc': '2.0', 'id': 1, 'method': 'initialize'},
|
|
|
+ )
|
|
|
+
|
|
|
+ self.assertIn('result', response)
|
|
|
+
|
|
|
+ def test_invalid_backend_result_emits_response_safety_failure(self):
|
|
|
+ reporter = RecordingReporter()
|
|
|
+ gateway = FakeGateway()
|
|
|
+ gateway.tool_result = []
|
|
|
+ handler = PublicMcpHttpHandler(
|
|
|
+ gateway,
|
|
|
+ context_parser=FakeParser(),
|
|
|
+ reporter=reporter,
|
|
|
+ )
|
|
|
+
|
|
|
+ response = handler.handle_json_rpc(
|
|
|
+ headers={'X-Gateway-Session': 'GWS_A'},
|
|
|
+ message={
|
|
|
+ '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'])
|
|
|
+ self.assertEqual('RESPONSE_SAFETY_REJECTED', event['event_code'])
|
|
|
+
|
|
|
+ def test_missing_tool_name_emits_protocol_validation_failure(self):
|
|
|
+ reporter = RecordingReporter()
|
|
|
+ handler = PublicMcpHttpHandler(
|
|
|
+ FakeGateway(),
|
|
|
+ context_parser=FakeParser(),
|
|
|
+ reporter=reporter,
|
|
|
+ )
|
|
|
+
|
|
|
+ response = handler.handle_json_rpc(
|
|
|
+ headers={'X-Gateway-Session': 'GWS_A'},
|
|
|
+ message={
|
|
|
+ 'jsonrpc': '2.0',
|
|
|
+ 'id': 1,
|
|
|
+ 'method': 'tools/call',
|
|
|
+ 'params': {'arguments': {}},
|
|
|
+ },
|
|
|
+ )
|
|
|
+
|
|
|
+ self.assertNotIn('result', response)
|
|
|
+ self.assertEqual(-32602, response['error']['code'])
|
|
|
+ self.assertTrue(response['error']['data']['request_id'].startswith('rq_http_'))
|
|
|
+ self.assertEqual(
|
|
|
+ 'failed',
|
|
|
+ next(
|
|
|
+ event for event in reporter.events
|
|
|
+ if event['stage'] == 'protocol_validation'
|
|
|
+ )['status'],
|
|
|
+ )
|
|
|
+
|
|
|
+ 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,
|
|
|
@@ -43,13 +264,118 @@ class PublicMcpHttpHandlerTest(unittest.TestCase):
|
|
|
'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_missing_session_on_tool_call_returns_error_content(self):
|
|
|
+ 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(
|
|
|
@@ -62,8 +388,40 @@ class PublicMcpHttpHandlerTest(unittest.TestCase):
|
|
|
},
|
|
|
)
|
|
|
|
|
|
+ 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.assertIn('这台设备的 Workbuddy 配置已失效,请重新生成配置', response['result']['content'][0]['text'])
|
|
|
+ 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(
|
|
|
@@ -73,6 +431,248 @@ class PublicMcpHttpHandlerTest(unittest.TestCase):
|
|
|
|
|
|
self.assertEqual('10.0.0.5', client_ip)
|
|
|
|
|
|
+
|
|
|
+class RateLimitTest(unittest.TestCase):
|
|
|
+ def test_rate_limit_helper_remains_usable_without_emitter(self):
|
|
|
+ handler = self._make_handler(max_requests=1)
|
|
|
+ self.assertIsNone(handler._check_rate_limit(
|
|
|
+ 'GWS_A:query_order', 'tools/call', 'rq_http_first'
|
|
|
+ ))
|
|
|
+
|
|
|
+ response = handler._check_rate_limit(
|
|
|
+ 'GWS_A:query_order', 'tools/call', 'rq_http_second'
|
|
|
+ )
|
|
|
+
|
|
|
+ self.assertEqual(-32029, response['error']['code'])
|
|
|
+
|
|
|
+ def test_rate_limit_rejection_emits_failed_diagnostic_event(self):
|
|
|
+ reporter = RecordingReporter()
|
|
|
+ limiter = SimpleRateLimiter(
|
|
|
+ max_requests=1,
|
|
|
+ window_seconds=60,
|
|
|
+ max_in_flight=1,
|
|
|
+ )
|
|
|
+ handler = PublicMcpHttpHandler(
|
|
|
+ FakeGateway(),
|
|
|
+ context_parser=FakeParser(),
|
|
|
+ rate_limiter=limiter,
|
|
|
+ reporter=reporter,
|
|
|
+ )
|
|
|
+ headers = {'X-Gateway-Session': 'GWS_A'}
|
|
|
+
|
|
|
+ handler.handle_json_rpc(headers, self._tools_call_msg())
|
|
|
+ reporter.events.clear()
|
|
|
+ handler.handle_json_rpc(headers, self._tools_call_msg())
|
|
|
+
|
|
|
+ event = next(
|
|
|
+ item for item in reporter.events if item['stage'] == 'rate_limit'
|
|
|
+ )
|
|
|
+ self.assertEqual('failed', event['status'])
|
|
|
+ self.assertEqual('RATE_LIMIT_EXCEEDED', event['event_code'])
|
|
|
+
|
|
|
+ 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_emits_response_write_failure(self):
|
|
|
+ reporter = RecordingReporter()
|
|
|
+ emitter = RequestDiagnosticEmitter(reporter, 'rq_http_write')
|
|
|
+ handler_class = create_http_handler(
|
|
|
+ FakeGateway(),
|
|
|
+ rate_limiter=None,
|
|
|
+ reporter=reporter,
|
|
|
+ )
|
|
|
+ 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()
|
|
|
+
|
|
|
+ handler._write_json(
|
|
|
+ {'jsonrpc': '2.0', 'id': 1, 'result': {}},
|
|
|
+ diagnostic_emitter=emitter,
|
|
|
+ )
|
|
|
+
|
|
|
+ self.assertEqual('response_write', reporter.events[-1]['stage'])
|
|
|
+ self.assertEqual('failed', reporter.events[-1]['status'])
|
|
|
+ self.assertTrue(
|
|
|
+ reporter.events[-1]['context']['client_disconnected']
|
|
|
+ )
|
|
|
+
|
|
|
+ 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])
|
|
|
+
|
|
|
+
|
|
|
+class RecordingReporter:
|
|
|
+ def __init__(self):
|
|
|
+ self.events = []
|
|
|
+
|
|
|
+ def report(self, event):
|
|
|
+ self.events.append(event)
|
|
|
+ return True
|
|
|
+
|
|
|
if __name__ == '__main__':
|
|
|
unittest.main()
|
|
|
-
|