test_public_server.py 28 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708
  1. import json
  2. import unittest
  3. from unittest.mock import Mock
  4. from public_server import PublicMcpHttpHandler, create_http_handler, extract_client_ip
  5. from services.diagnostic_event import RequestDiagnosticEmitter
  6. from utils.rate_limiter import SimpleRateLimiter
  7. from constants import DEVICE_INVALID_MESSAGE
  8. class FakeContext:
  9. def __init__(self, gateway_session_id):
  10. self.gateway_session_id = gateway_session_id
  11. def has_session(self):
  12. return bool(self.gateway_session_id)
  13. class FakeParser:
  14. def parse(self, headers):
  15. return FakeContext(headers.get('X-Gateway-Session', ''))
  16. class FakeGateway:
  17. def __init__(self):
  18. self.calls = []
  19. self.list_calls = []
  20. self.tool_result = {'code': 'MCP_0000', 'data': {'ok': True}}
  21. def registered_tool_names(self):
  22. return ('query_order', 'query_track')
  23. def list_tools(self, gateway_session_id, request_id=''):
  24. self.list_calls.append((gateway_session_id, request_id))
  25. return [{'name': 'query_track', 'description': 'query track', 'input_schema': {'type': 'object'}}]
  26. def call_tool(
  27. self,
  28. gateway_session_id,
  29. name,
  30. arguments=None,
  31. request_id='',
  32. client_ip='',
  33. diagnostic_emitter=None,
  34. ):
  35. self.calls.append((gateway_session_id, name, arguments, request_id, client_ip))
  36. return self.tool_result
  37. class PublicMcpHttpHandlerTest(unittest.TestCase):
  38. def test_initialize_emits_successful_protocol_validation(self):
  39. reporter = RecordingReporter()
  40. handler = PublicMcpHttpHandler(FakeGateway(), reporter=reporter)
  41. response = handler.handle_json_rpc(
  42. headers={},
  43. message={'jsonrpc': '2.0', 'id': 1, 'method': 'initialize'},
  44. )
  45. self.assertIn('result', response)
  46. self.assertEqual(
  47. ['request_ingress', 'protocol_validation'],
  48. [event['stage'] for event in reporter.events],
  49. )
  50. self.assertEqual('succeeded', reporter.events[-1]['status'])
  51. def test_missing_session_emits_safe_ingress_and_session_failure(self):
  52. reporter = RecordingReporter()
  53. handler = PublicMcpHttpHandler(
  54. FakeGateway(),
  55. context_parser=FakeParser(),
  56. reporter=reporter,
  57. )
  58. response = handler.handle_json_rpc(
  59. headers={},
  60. message={
  61. 'jsonrpc': '2.0',
  62. 'id': 1,
  63. 'method': 'tools/call',
  64. 'params': {'name': 'query_order', 'arguments': {}},
  65. },
  66. )
  67. self.assertEqual(-32001, response['error']['code'])
  68. self.assertEqual(
  69. ['request_ingress', 'protocol_validation', 'gateway_session'],
  70. [event['stage'] for event in reporter.events],
  71. )
  72. self.assertEqual('failed', reporter.events[-1]['status'])
  73. def test_reporter_failure_does_not_change_successful_response(self):
  74. class BrokenReporter:
  75. def report(self, _event):
  76. raise RuntimeError('support unavailable')
  77. handler = PublicMcpHttpHandler(
  78. FakeGateway(),
  79. context_parser=FakeParser(),
  80. reporter=BrokenReporter(),
  81. )
  82. response = handler.handle_json_rpc(
  83. headers={'X-Gateway-Session': 'GWS_A'},
  84. message={'jsonrpc': '2.0', 'id': 1, 'method': 'initialize'},
  85. )
  86. self.assertIn('result', response)
  87. def test_invalid_backend_result_emits_response_safety_failure(self):
  88. reporter = RecordingReporter()
  89. gateway = FakeGateway()
  90. gateway.tool_result = []
  91. handler = PublicMcpHttpHandler(
  92. gateway,
  93. context_parser=FakeParser(),
  94. reporter=reporter,
  95. )
  96. response = handler.handle_json_rpc(
  97. headers={'X-Gateway-Session': 'GWS_A'},
  98. message={
  99. 'jsonrpc': '2.0',
  100. 'id': 1,
  101. 'method': 'tools/call',
  102. 'params': {'name': 'query_order', 'arguments': {}},
  103. },
  104. )
  105. self.assertTrue(response['result']['isError'])
  106. event = next(
  107. item for item in reporter.events
  108. if item['stage'] == 'response_safety'
  109. )
  110. self.assertEqual('failed', event['status'])
  111. self.assertEqual('RESPONSE_SAFETY_REJECTED', event['event_code'])
  112. def test_missing_tool_name_emits_protocol_validation_failure(self):
  113. reporter = RecordingReporter()
  114. handler = PublicMcpHttpHandler(
  115. FakeGateway(),
  116. context_parser=FakeParser(),
  117. reporter=reporter,
  118. )
  119. response = handler.handle_json_rpc(
  120. headers={'X-Gateway-Session': 'GWS_A'},
  121. message={
  122. 'jsonrpc': '2.0',
  123. 'id': 1,
  124. 'method': 'tools/call',
  125. 'params': {'arguments': {}},
  126. },
  127. )
  128. self.assertNotIn('result', response)
  129. self.assertEqual(-32602, response['error']['code'])
  130. self.assertTrue(response['error']['data']['request_id'].startswith('rq_http_'))
  131. self.assertEqual(
  132. 'failed',
  133. next(
  134. event for event in reporter.events
  135. if event['stage'] == 'protocol_validation'
  136. )['status'],
  137. )
  138. def test_constructor_uses_registered_names_without_loading_dynamic_list(self):
  139. gateway = FakeGateway()
  140. PublicMcpHttpHandler(gateway, context_parser=FakeParser())
  141. self.assertEqual([], gateway.list_calls)
  142. def test_handle_tools_list_passes_session_and_returns_dynamic_tools(self):
  143. gateway = FakeGateway()
  144. handler = PublicMcpHttpHandler(gateway, context_parser=FakeParser())
  145. response = handler.handle_json_rpc(
  146. headers={'X-Gateway-Session': 'GWS_A'},
  147. message={
  148. 'jsonrpc': '2.0',
  149. 'id': 2,
  150. 'method': 'tools/list',
  151. 'params': {},
  152. },
  153. client_ip='10.0.0.5',
  154. )
  155. self.assertEqual('GWS_A', gateway.list_calls[0][0])
  156. self.assertTrue(gateway.list_calls[0][1].startswith('rq_http_'))
  157. self.assertEqual(
  158. ['query_track'],
  159. [tool['name'] for tool in response['result']['tools']],
  160. )
  161. self.assertIn('inputSchema', response['result']['tools'][0])
  162. def test_missing_session_on_tools_list_returns_protocol_error(self):
  163. gateway = FakeGateway()
  164. handler = PublicMcpHttpHandler(gateway, context_parser=FakeParser())
  165. with self.assertLogs('public_server', level='WARNING') as logs:
  166. response = handler.handle_json_rpc(
  167. headers={},
  168. message={
  169. 'jsonrpc': '2.0',
  170. 'id': 3,
  171. 'method': 'tools/list',
  172. 'params': {},
  173. },
  174. client_ip='10.0.0.5',
  175. )
  176. self.assertEqual(-32001, response['error']['code'])
  177. self.assertIn('Workbuddy', response['error']['message'])
  178. self.assertTrue(response['error']['data']['request_id'].startswith('rq_http_'))
  179. self.assertEqual([], gateway.list_calls)
  180. record = logs.records[0]
  181. self.assertEqual(3, record.jsonrpc_id)
  182. self.assertEqual(response['error']['data']['request_id'], record.request_id)
  183. self.assertEqual('tools/list', record.protocol_method)
  184. self.assertEqual(-32001, record.protocol_code)
  185. self.assertEqual('GATEWAY_SESSION_NOT_FOUND', record.diagnostic_reason)
  186. def test_client_request_id_is_logged_only_as_safe_hash(self):
  187. gateway = FakeGateway()
  188. handler = PublicMcpHttpHandler(gateway, context_parser=FakeParser())
  189. value = handler._client_request_id({
  190. 'X-Request-Id': 'client-id\nAuthorization: secret',
  191. })
  192. self.assertRegex(value, r'^[0-9a-f]{16}$')
  193. self.assertNotIn('client-id', value)
  194. def test_tools_list_internal_exception_is_sanitized(self):
  195. class ExplodingGateway(FakeGateway):
  196. def list_tools(self, gateway_session_id, request_id=''):
  197. raise RuntimeError('database password leaked')
  198. handler = PublicMcpHttpHandler(ExplodingGateway(), context_parser=FakeParser())
  199. response = handler.handle_json_rpc(
  200. headers={'X-Gateway-Session': 'GWS_A'},
  201. message={'jsonrpc': '2.0', 'id': 31, 'method': 'tools/list', 'params': {}},
  202. client_ip='10.0.0.5',
  203. )
  204. self.assertEqual(-32000, response['error']['code'])
  205. self.assertEqual('Gateway request failed. Please try again later.', response['error']['message'])
  206. self.assertNotIn('password', json.dumps(response))
  207. def test_tools_list_redis_miss_returns_device_invalid(self):
  208. class MissingSessionGateway(FakeGateway):
  209. def list_tools(self, gateway_session_id, request_id=''):
  210. raise RuntimeError(DEVICE_INVALID_MESSAGE)
  211. reporter = RecordingReporter()
  212. handler = PublicMcpHttpHandler(
  213. MissingSessionGateway(),
  214. context_parser=FakeParser(),
  215. reporter=reporter,
  216. )
  217. with self.assertLogs('public_server', level='WARNING') as logs:
  218. response = handler.handle_json_rpc(
  219. headers={'X-Gateway-Session': 'GWS_A'},
  220. message={'jsonrpc': '2.0', 'id': 32, 'method': 'tools/list', 'params': {}},
  221. client_ip='10.0.0.5',
  222. )
  223. self.assertEqual(-32001, response['error']['code'])
  224. self.assertEqual(DEVICE_INVALID_MESSAGE, response['error']['message'])
  225. self.assertTrue(response['error']['data']['request_id'].startswith('rq_http_'))
  226. event = next(item for item in reporter.events if item['stage'] == 'gateway_session')
  227. self.assertEqual('failed', event['status'])
  228. self.assertEqual('GATEWAY_SESSION_NOT_FOUND', event['event_code'])
  229. record = logs.records[0]
  230. self.assertEqual(-32001, record.protocol_code)
  231. self.assertEqual('GATEWAY_SESSION_NOT_FOUND', record.diagnostic_reason)
  232. def test_handle_tools_call_passes_gateway_session_to_public_gateway(self):
  233. gateway = FakeGateway()
  234. handler = PublicMcpHttpHandler(gateway, context_parser=FakeParser())
  235. response = handler.handle_json_rpc(
  236. headers={
  237. 'X-Gateway-Session': 'GWS_A',
  238. 'X-Request-Id': 'rq_public_incoming',
  239. },
  240. message={
  241. 'jsonrpc': '2.0',
  242. 'id': 1,
  243. 'method': 'tools/call',
  244. 'params': {
  245. 'name': 'query_order',
  246. 'arguments': {'keyword': 'USC'},
  247. },
  248. },
  249. client_ip='10.0.0.5'
  250. )
  251. self.assertEqual(False, response['result']['isError'])
  252. self.assertEqual('GWS_A', gateway.calls[0][0])
  253. self.assertEqual('query_order', gateway.calls[0][1])
  254. self.assertTrue(gateway.calls[0][3].startswith('rq_http_'))
  255. self.assertNotEqual('rq_public_incoming', gateway.calls[0][3])
  256. self.assertEqual('10.0.0.5', gateway.calls[0][4])
  257. def test_reused_client_request_id_gets_unique_server_trace_ids(self):
  258. gateway = FakeGateway()
  259. handler = PublicMcpHttpHandler(gateway, context_parser=FakeParser())
  260. headers = {
  261. 'X-Gateway-Session': 'GWS_A',
  262. 'X-Request-Id': 'client-retry-id',
  263. }
  264. message = {
  265. 'jsonrpc': '2.0',
  266. 'id': 1,
  267. 'method': 'tools/call',
  268. 'params': {'name': 'query_track', 'arguments': {}},
  269. }
  270. handler.handle_json_rpc(headers, message, client_ip='10.0.0.5')
  271. handler.handle_json_rpc(headers, message, client_ip='10.0.0.5')
  272. first_request_id = gateway.calls[0][3]
  273. second_request_id = gateway.calls[1][3]
  274. self.assertTrue(first_request_id.startswith('rq_http_'))
  275. self.assertTrue(second_request_id.startswith('rq_http_'))
  276. self.assertNotEqual(first_request_id, second_request_id)
  277. def test_backend_business_error_returns_tool_error_result(self):
  278. gateway = FakeGateway()
  279. gateway.tool_result = {
  280. 'code': 'MCP_1401',
  281. 'msg': 'invalid exact order query conditions',
  282. 'data': [],
  283. 'meta': {'request_id': 'rq_backend'},
  284. }
  285. handler = PublicMcpHttpHandler(gateway, context_parser=FakeParser())
  286. response = handler.handle_json_rpc(
  287. headers={'X-Gateway-Session': 'GWS_A'},
  288. message={
  289. 'jsonrpc': '2.0',
  290. 'id': 7,
  291. 'method': 'tools/call',
  292. 'params': {
  293. 'name': 'query_order',
  294. 'arguments': {'keyword': 'USC'},
  295. },
  296. },
  297. client_ip='10.0.0.5',
  298. )
  299. self.assertNotIn('error', response)
  300. self.assertTrue(response['result']['isError'])
  301. self.assertIn('MCP_1401', response['result']['content'][0]['text'])
  302. self.assertEqual(
  303. {
  304. 'code': 'MCP_1401',
  305. 'msg': 'invalid exact order query conditions',
  306. 'meta': {'request_id': 'rq_backend'},
  307. },
  308. response['result']['structuredContent'],
  309. )
  310. def test_query_track_success_uses_safe_public_output(self):
  311. gateway = FakeGateway()
  312. gateway.tool_result = {
  313. 'code': 'MCP_0000',
  314. 'data': {
  315. 'summary': '共 1 条轨迹',
  316. 'columns': [
  317. {'key': 'status', 'name': '轨迹节点'},
  318. {'key': 'tracking_number', 'name': '快递单号'},
  319. ],
  320. 'records': [{
  321. 'status': '已发货',
  322. 'tracking_number': 'TN-PUBLIC',
  323. 'time_zone': '8.00',
  324. }],
  325. },
  326. 'meta': {'request_id': 'rq_public', 'page': 1, 'limit': 5},
  327. }
  328. handler = PublicMcpHttpHandler(gateway, context_parser=FakeParser())
  329. response = handler.handle_json_rpc(
  330. headers={'X-Gateway-Session': 'GWS_A'},
  331. message={
  332. 'jsonrpc': '2.0',
  333. 'id': 8,
  334. 'method': 'tools/call',
  335. 'params': {
  336. 'name': 'query_track',
  337. 'arguments': {'tracking_number': 'TN-PUBLIC'},
  338. },
  339. },
  340. client_ip='10.0.0.5',
  341. )
  342. serialized = json.dumps(response['result'], ensure_ascii=False)
  343. self.assertFalse(response['result']['isError'])
  344. self.assertNotIn('tracking_number', serialized)
  345. self.assertNotIn('time_zone', serialized)
  346. self.assertIn('快递单号', serialized)
  347. self.assertIn('TN-PUBLIC', serialized)
  348. self.assertEqual('rq_public', response['result']['_meta']['request_id'])
  349. def test_missing_session_on_tool_call_returns_device_protocol_error(self):
  350. gateway = FakeGateway()
  351. handler = PublicMcpHttpHandler(gateway, context_parser=FakeParser())
  352. response = handler.handle_json_rpc(
  353. headers={},
  354. message={
  355. 'jsonrpc': '2.0',
  356. 'id': 1,
  357. 'method': 'tools/call',
  358. 'params': {'name': 'query_order', 'arguments': {'keyword': 'USC'}},
  359. },
  360. )
  361. self.assertEqual(-32001, response['error']['code'])
  362. self.assertIn('Workbuddy', response['error']['message'])
  363. self.assertEqual([], gateway.calls)
  364. def test_non_object_tool_params_fail_safely(self):
  365. gateway = FakeGateway()
  366. handler = PublicMcpHttpHandler(gateway, context_parser=FakeParser())
  367. with self.assertLogs('public_server', level='ERROR') as logs:
  368. response = handler.handle_json_rpc(
  369. headers={'X-Gateway-Session': 'GWS_A'},
  370. message={
  371. 'jsonrpc': '2.0',
  372. 'id': 9,
  373. 'method': 'tools/call',
  374. 'params': 'not-an-object',
  375. },
  376. client_ip='10.0.0.5',
  377. )
  378. self.assertTrue(response['result']['isError'])
  379. self.assertEqual(
  380. '工具返回格式异常',
  381. response['result']['structuredContent']['message'],
  382. )
  383. self.assertEqual([], gateway.calls)
  384. record = logs.records[0]
  385. self.assertEqual(9, record.jsonrpc_id)
  386. self.assertTrue(record.request_id.startswith('rq_http_'))
  387. self.assertEqual('', record.tool_code)
  388. self.assertEqual('MCP_9001', record.response_code)
  389. self.assertEqual('PARAM_VALIDATION_FAILED', record.diagnostic_reason)
  390. self.assertEqual('ValueError', record.exception_class)
  391. self.assertEqual(record.request_id, response['result']['_meta']['request_id'])
  392. def test_extract_client_ip_ignores_spoofable_forwarded_for_header(self):
  393. client_ip = extract_client_ip(
  394. {'X-Forwarded-For': '203.0.113.9'},
  395. ('10.0.0.5', 54321),
  396. )
  397. self.assertEqual('10.0.0.5', client_ip)
  398. class RateLimitTest(unittest.TestCase):
  399. def test_rate_limit_helper_remains_usable_without_emitter(self):
  400. handler = self._make_handler(max_requests=1)
  401. self.assertIsNone(handler._check_rate_limit(
  402. 'GWS_A:query_order', 'tools/call', 'rq_http_first'
  403. ))
  404. response = handler._check_rate_limit(
  405. 'GWS_A:query_order', 'tools/call', 'rq_http_second'
  406. )
  407. self.assertEqual(-32029, response['error']['code'])
  408. def test_rate_limit_rejection_emits_failed_diagnostic_event(self):
  409. reporter = RecordingReporter()
  410. limiter = SimpleRateLimiter(
  411. max_requests=1,
  412. window_seconds=60,
  413. max_in_flight=1,
  414. )
  415. handler = PublicMcpHttpHandler(
  416. FakeGateway(),
  417. context_parser=FakeParser(),
  418. rate_limiter=limiter,
  419. reporter=reporter,
  420. )
  421. headers = {'X-Gateway-Session': 'GWS_A'}
  422. handler.handle_json_rpc(headers, self._tools_call_msg())
  423. reporter.events.clear()
  424. handler.handle_json_rpc(headers, self._tools_call_msg())
  425. event = next(
  426. item for item in reporter.events if item['stage'] == 'rate_limit'
  427. )
  428. self.assertEqual('failed', event['status'])
  429. self.assertEqual('RATE_LIMIT_EXCEEDED', event['event_code'])
  430. def _make_handler(self, max_requests=2, max_in_flight=2):
  431. gateway = FakeGateway()
  432. limiter = SimpleRateLimiter(
  433. max_requests=max_requests,
  434. window_seconds=60,
  435. max_in_flight=max_in_flight,
  436. )
  437. return PublicMcpHttpHandler(gateway, context_parser=FakeParser(), rate_limiter=limiter)
  438. def _tools_call_msg(self, tool_name='query_order'):
  439. return {
  440. 'jsonrpc': '2.0', 'id': 1,
  441. 'method': 'tools/call',
  442. 'params': {'name': tool_name, 'arguments': {'keyword': 'test'}},
  443. }
  444. # --- tools/call per-tool独立限流 ---
  445. def test_known_tool_rate_limited_after_quota_exhausted(self):
  446. handler = self._make_handler(max_requests=1)
  447. handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg(), client_ip='1.2.3.4')
  448. with self.assertLogs('public_server', level='WARNING') as logs:
  449. response = handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg(), client_ip='1.2.3.4')
  450. self.assertIn('error', response)
  451. self.assertEqual(-32029, response['error']['code'])
  452. self.assertIn('Rate limit', response['error']['message'])
  453. record = logs.records[0]
  454. self.assertEqual(1, record.jsonrpc_id)
  455. self.assertEqual('query_order', record.tool_code)
  456. self.assertEqual(-32029, record.protocol_code)
  457. self.assertEqual('RATE_LIMIT_EXCEEDED', record.diagnostic_reason)
  458. self.assertEqual(response['error']['data']['request_id'], record.request_id)
  459. def test_two_known_tools_have_independent_quotas(self):
  460. # query_order 限流不影响 query_track
  461. handler = self._make_handler(max_requests=1)
  462. handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg('query_order'), client_ip='1.2.3.4')
  463. # query_order quota exhausted, query_track should still work
  464. response = handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg('query_track'), client_ip='1.2.3.4')
  465. self.assertNotIn('error', response)
  466. def test_unknown_tool_name_uses_shared_session_bucket_not_new_bucket(self):
  467. # 未知工具名应归入 session_id 桶,同一会话的不同未知工具共享同一配额,不能无限新建桶
  468. handler = self._make_handler(max_requests=1)
  469. # 先用 session 桶打一次(用未知工具名触发)
  470. handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg('fake_tool_1'), client_ip='1.2.3.4')
  471. # 再用同会话的另一个未知工具名,应该命中同一个 session 桶,被限流
  472. response = handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg('fake_tool_2'), client_ip='1.2.3.4')
  473. self.assertIn('error', response)
  474. self.assertIn('Rate limit', response['error']['message'])
  475. def test_different_sessions_have_independent_quotas(self):
  476. # 不同会话(不同员工)即使来自同一IP,也拥有独立的配额
  477. handler = self._make_handler(max_requests=1)
  478. handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg(), client_ip='1.1.1.1')
  479. # 不同 session,即使同一 IP,quota 独立 → 应该允许
  480. response = handler.handle_json_rpc({'X-Gateway-Session': 'GWS_B'}, self._tools_call_msg(), client_ip='1.1.1.1')
  481. self.assertNotIn('error', response)
  482. def test_same_session_from_different_ips_shares_quota(self):
  483. # 同一会话从不同IP发来(如移动网络切换),仍共享同一会话配额
  484. handler = self._make_handler(max_requests=1)
  485. handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg(), client_ip='1.1.1.1')
  486. # 同 session,不同 IP → 命中同一 session 桶,被限流
  487. response = handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg(), client_ip='2.2.2.2')
  488. self.assertIn('error', response)
  489. self.assertIn('Rate limit', response['error']['message'])
  490. # --- initialize 和 tools/list 不受限流 ---
  491. def test_initialize_is_never_rate_limited(self):
  492. # initialize 是握手协议,无论请求多少次都不应被限流
  493. handler = self._make_handler(max_requests=1)
  494. for i in range(5):
  495. response = handler.handle_json_rpc({}, {'jsonrpc': '2.0', 'id': i, 'method': 'initialize'}, client_ip='1.2.3.4')
  496. self.assertNotIn('error', response, f'initialize should never be rate limited (attempt {i})')
  497. self.assertIn('protocolVersion', response['result'])
  498. def test_tools_list_is_never_rate_limited(self):
  499. handler = self._make_handler(max_requests=1)
  500. headers = {'X-Gateway-Session': 'GWS_A'}
  501. for request_id in range(1, 4):
  502. response = handler.handle_json_rpc(
  503. headers,
  504. {'jsonrpc': '2.0', 'id': request_id, 'method': 'tools/list'},
  505. client_ip='1.2.3.4',
  506. )
  507. self.assertNotIn('error', response)
  508. def test_tools_list_requests_do_not_consume_tools_call_quota(self):
  509. # tools/list 不进入限流器,因此不会消耗 tools/call 配额
  510. handler = self._make_handler(max_requests=1)
  511. handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, {'jsonrpc': '2.0', 'id': 1, 'method': 'tools/list'}, client_ip='1.2.3.4')
  512. # tools/call 应仍可正常执行(使用 session_id:tool_name bucket)
  513. response = handler.handle_json_rpc(
  514. {'X-Gateway-Session': 'GWS_A'},
  515. {'jsonrpc': '2.0', 'id': 2, 'method': 'tools/call',
  516. 'params': {'name': 'query_order', 'arguments': {'keyword': 'test'}}},
  517. client_ip='1.2.3.4',
  518. )
  519. self.assertNotIn('error', response)
  520. def test_completed_tool_call_releases_in_flight_slot_immediately(self):
  521. handler = self._make_handler(max_requests=0, max_in_flight=1)
  522. headers = {'X-Gateway-Session': 'GWS_A'}
  523. for request_id in range(1, 6):
  524. message = self._tools_call_msg('query_order')
  525. message['id'] = request_id
  526. response = handler.handle_json_rpc(
  527. headers,
  528. message,
  529. client_ip='1.2.3.4',
  530. )
  531. self.assertNotIn('error', response)
  532. def test_exhausted_in_flight_slots_return_rate_limit_error(self):
  533. handler = self._make_handler(max_requests=0, max_in_flight=2)
  534. key = 'GWS_A:query_order'
  535. self.assertTrue(handler.rate_limiter.try_acquire(key))
  536. self.assertTrue(handler.rate_limiter.try_acquire(key))
  537. response = handler.handle_json_rpc(
  538. {'X-Gateway-Session': 'GWS_A'},
  539. self._tools_call_msg('query_order'),
  540. client_ip='1.2.3.4',
  541. )
  542. self.assertEqual(-32029, response['error']['code'])
  543. self.assertIn('in progress', response['error']['message'])
  544. def test_failed_tool_call_releases_in_flight_slot(self):
  545. handler = self._make_handler(max_requests=0, max_in_flight=1)
  546. original_call = handler.gateway_app.call_tool
  547. def fail(*args, **kwargs):
  548. raise RuntimeError('backend failed')
  549. handler.gateway_app.call_tool = fail
  550. handler.handle_json_rpc(
  551. {'X-Gateway-Session': 'GWS_A'},
  552. self._tools_call_msg('query_order'),
  553. client_ip='1.2.3.4',
  554. )
  555. handler.gateway_app.call_tool = original_call
  556. response = handler.handle_json_rpc(
  557. {'X-Gateway-Session': 'GWS_A'},
  558. self._tools_call_msg('query_order'),
  559. client_ip='1.2.3.4',
  560. )
  561. self.assertNotIn('error', response)
  562. class HttpDisconnectTest(unittest.TestCase):
  563. def test_broken_pipe_emits_response_write_failure(self):
  564. reporter = RecordingReporter()
  565. emitter = RequestDiagnosticEmitter(reporter, 'rq_http_write')
  566. handler_class = create_http_handler(
  567. FakeGateway(),
  568. rate_limiter=None,
  569. reporter=reporter,
  570. )
  571. handler = object.__new__(handler_class)
  572. handler.send_response = Mock()
  573. handler.send_header = Mock()
  574. handler.end_headers = Mock()
  575. handler.wfile = Mock()
  576. handler.wfile.write.side_effect = BrokenPipeError()
  577. handler._write_json(
  578. {'jsonrpc': '2.0', 'id': 1, 'result': {}},
  579. diagnostic_emitter=emitter,
  580. )
  581. self.assertEqual('response_write', reporter.events[-1]['stage'])
  582. self.assertEqual('failed', reporter.events[-1]['status'])
  583. self.assertTrue(
  584. reporter.events[-1]['context']['client_disconnected']
  585. )
  586. def test_broken_pipe_while_writing_response_is_handled(self):
  587. handler_class = create_http_handler(FakeGateway(), rate_limiter=None)
  588. handler = object.__new__(handler_class)
  589. handler.send_response = Mock()
  590. handler.send_header = Mock()
  591. handler.end_headers = Mock()
  592. handler.wfile = Mock()
  593. handler.wfile.write.side_effect = BrokenPipeError()
  594. with self.assertLogs('public_server', level='INFO') as logs:
  595. handler._write_json({'jsonrpc': '2.0', 'id': 1, 'result': {}})
  596. self.assertIn('client disconnected before response', logs.output[0])
  597. class RecordingReporter:
  598. def __init__(self):
  599. self.events = []
  600. def report(self, event):
  601. self.events.append(event)
  602. return True
  603. if __name__ == '__main__':
  604. unittest.main()