test_public_server.py 26 KB

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