test_public_server.py 20 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482
  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 utils.rate_limiter import SimpleRateLimiter
  6. class FakeContext:
  7. def __init__(self, gateway_session_id):
  8. self.gateway_session_id = gateway_session_id
  9. def has_session(self):
  10. return bool(self.gateway_session_id)
  11. class FakeParser:
  12. def parse(self, headers):
  13. return FakeContext(headers.get('X-Gateway-Session', ''))
  14. class FakeGateway:
  15. def __init__(self):
  16. self.calls = []
  17. self.list_calls = []
  18. self.tool_result = {'code': 'MCP_0000', 'data': {'ok': True}}
  19. def registered_tool_names(self):
  20. return ('query_order', 'query_track')
  21. def list_tools(self, gateway_session_id, request_id=''):
  22. self.list_calls.append((gateway_session_id, request_id))
  23. return [{'name': 'query_track', 'description': 'query track', 'input_schema': {'type': 'object'}}]
  24. def call_tool(self, gateway_session_id, name, arguments=None, request_id='', client_ip=''):
  25. self.calls.append((gateway_session_id, name, arguments, request_id, client_ip))
  26. return self.tool_result
  27. class PublicMcpHttpHandlerTest(unittest.TestCase):
  28. def test_constructor_uses_registered_names_without_loading_dynamic_list(self):
  29. gateway = FakeGateway()
  30. PublicMcpHttpHandler(gateway, context_parser=FakeParser())
  31. self.assertEqual([], gateway.list_calls)
  32. def test_handle_tools_list_passes_session_and_returns_dynamic_tools(self):
  33. gateway = FakeGateway()
  34. handler = PublicMcpHttpHandler(gateway, context_parser=FakeParser())
  35. response = handler.handle_json_rpc(
  36. headers={'X-Gateway-Session': 'GWS_A'},
  37. message={
  38. 'jsonrpc': '2.0',
  39. 'id': 2,
  40. 'method': 'tools/list',
  41. 'params': {},
  42. },
  43. client_ip='10.0.0.5',
  44. )
  45. self.assertEqual('GWS_A', gateway.list_calls[0][0])
  46. self.assertTrue(gateway.list_calls[0][1].startswith('rq_http_'))
  47. self.assertEqual(
  48. ['query_track'],
  49. [tool['name'] for tool in response['result']['tools']],
  50. )
  51. self.assertIn('inputSchema', response['result']['tools'][0])
  52. def test_missing_session_on_tools_list_returns_protocol_error(self):
  53. gateway = FakeGateway()
  54. handler = PublicMcpHttpHandler(gateway, context_parser=FakeParser())
  55. with self.assertLogs('public_server', level='WARNING') as logs:
  56. response = handler.handle_json_rpc(
  57. headers={},
  58. message={
  59. 'jsonrpc': '2.0',
  60. 'id': 3,
  61. 'method': 'tools/list',
  62. 'params': {},
  63. },
  64. client_ip='10.0.0.5',
  65. )
  66. self.assertEqual(-32001, response['error']['code'])
  67. self.assertIn('Workbuddy', response['error']['message'])
  68. self.assertTrue(response['error']['data']['request_id'].startswith('rq_http_'))
  69. self.assertEqual([], gateway.list_calls)
  70. record = logs.records[0]
  71. self.assertEqual(3, record.jsonrpc_id)
  72. self.assertEqual(response['error']['data']['request_id'], record.request_id)
  73. self.assertEqual('tools/list', record.protocol_method)
  74. self.assertEqual(-32001, record.protocol_code)
  75. self.assertEqual('GATEWAY_SESSION_NOT_FOUND', record.diagnostic_reason)
  76. def test_client_request_id_is_logged_only_as_safe_hash(self):
  77. gateway = FakeGateway()
  78. handler = PublicMcpHttpHandler(gateway, context_parser=FakeParser())
  79. value = handler._client_request_id({
  80. 'X-Request-Id': 'client-id\nAuthorization: secret',
  81. })
  82. self.assertRegex(value, r'^[0-9a-f]{16}$')
  83. self.assertNotIn('client-id', value)
  84. def test_tools_list_internal_exception_is_sanitized(self):
  85. class ExplodingGateway(FakeGateway):
  86. def list_tools(self, gateway_session_id, request_id=''):
  87. raise RuntimeError('database password leaked')
  88. handler = PublicMcpHttpHandler(ExplodingGateway(), context_parser=FakeParser())
  89. response = handler.handle_json_rpc(
  90. headers={'X-Gateway-Session': 'GWS_A'},
  91. message={'jsonrpc': '2.0', 'id': 31, 'method': 'tools/list', 'params': {}},
  92. client_ip='10.0.0.5',
  93. )
  94. self.assertEqual(-32000, response['error']['code'])
  95. self.assertEqual('Gateway request failed. Please try again later.', response['error']['message'])
  96. self.assertNotIn('password', json.dumps(response))
  97. def test_handle_tools_call_passes_gateway_session_to_public_gateway(self):
  98. gateway = FakeGateway()
  99. handler = PublicMcpHttpHandler(gateway, context_parser=FakeParser())
  100. response = handler.handle_json_rpc(
  101. headers={
  102. 'X-Gateway-Session': 'GWS_A',
  103. 'X-Request-Id': 'rq_public_incoming',
  104. },
  105. message={
  106. 'jsonrpc': '2.0',
  107. 'id': 1,
  108. 'method': 'tools/call',
  109. 'params': {
  110. 'name': 'query_order',
  111. 'arguments': {'keyword': 'USC'},
  112. },
  113. },
  114. client_ip='10.0.0.5'
  115. )
  116. self.assertEqual(False, response['result']['isError'])
  117. self.assertEqual('GWS_A', gateway.calls[0][0])
  118. self.assertEqual('query_order', gateway.calls[0][1])
  119. self.assertTrue(gateway.calls[0][3].startswith('rq_http_'))
  120. self.assertNotEqual('rq_public_incoming', gateway.calls[0][3])
  121. self.assertEqual('10.0.0.5', gateway.calls[0][4])
  122. def test_reused_client_request_id_gets_unique_server_trace_ids(self):
  123. gateway = FakeGateway()
  124. handler = PublicMcpHttpHandler(gateway, context_parser=FakeParser())
  125. headers = {
  126. 'X-Gateway-Session': 'GWS_A',
  127. 'X-Request-Id': 'client-retry-id',
  128. }
  129. message = {
  130. 'jsonrpc': '2.0',
  131. 'id': 1,
  132. 'method': 'tools/call',
  133. 'params': {'name': 'query_track', 'arguments': {}},
  134. }
  135. handler.handle_json_rpc(headers, message, client_ip='10.0.0.5')
  136. handler.handle_json_rpc(headers, message, client_ip='10.0.0.5')
  137. first_request_id = gateway.calls[0][3]
  138. second_request_id = gateway.calls[1][3]
  139. self.assertTrue(first_request_id.startswith('rq_http_'))
  140. self.assertTrue(second_request_id.startswith('rq_http_'))
  141. self.assertNotEqual(first_request_id, second_request_id)
  142. def test_backend_business_error_returns_tool_error_result(self):
  143. gateway = FakeGateway()
  144. gateway.tool_result = {
  145. 'code': 'MCP_1401',
  146. 'msg': 'invalid exact order query conditions',
  147. 'data': [],
  148. 'meta': {'request_id': 'rq_backend'},
  149. }
  150. handler = PublicMcpHttpHandler(gateway, context_parser=FakeParser())
  151. response = handler.handle_json_rpc(
  152. headers={'X-Gateway-Session': 'GWS_A'},
  153. message={
  154. 'jsonrpc': '2.0',
  155. 'id': 7,
  156. 'method': 'tools/call',
  157. 'params': {
  158. 'name': 'query_order',
  159. 'arguments': {'keyword': 'USC'},
  160. },
  161. },
  162. client_ip='10.0.0.5',
  163. )
  164. self.assertNotIn('error', response)
  165. self.assertTrue(response['result']['isError'])
  166. self.assertIn('MCP_1401', response['result']['content'][0]['text'])
  167. self.assertEqual(
  168. {
  169. 'code': 'MCP_1401',
  170. 'msg': 'invalid exact order query conditions',
  171. 'meta': {'request_id': 'rq_backend'},
  172. },
  173. response['result']['structuredContent'],
  174. )
  175. def test_query_track_success_uses_safe_public_output(self):
  176. gateway = FakeGateway()
  177. gateway.tool_result = {
  178. 'code': 'MCP_0000',
  179. 'data': {
  180. 'summary': '共 1 条轨迹',
  181. 'columns': [
  182. {'key': 'status', 'name': '轨迹节点'},
  183. {'key': 'tracking_number', 'name': '快递单号'},
  184. ],
  185. 'records': [{
  186. 'status': '已发货',
  187. 'tracking_number': 'TN-PUBLIC',
  188. 'time_zone': '8.00',
  189. }],
  190. },
  191. 'meta': {'request_id': 'rq_public', 'page': 1, 'limit': 5},
  192. }
  193. handler = PublicMcpHttpHandler(gateway, context_parser=FakeParser())
  194. response = handler.handle_json_rpc(
  195. headers={'X-Gateway-Session': 'GWS_A'},
  196. message={
  197. 'jsonrpc': '2.0',
  198. 'id': 8,
  199. 'method': 'tools/call',
  200. 'params': {
  201. 'name': 'query_track',
  202. 'arguments': {'tracking_number': 'TN-PUBLIC'},
  203. },
  204. },
  205. client_ip='10.0.0.5',
  206. )
  207. serialized = json.dumps(response['result'], ensure_ascii=False)
  208. self.assertFalse(response['result']['isError'])
  209. self.assertNotIn('tracking_number', serialized)
  210. self.assertNotIn('time_zone', serialized)
  211. self.assertIn('快递单号', serialized)
  212. self.assertIn('TN-PUBLIC', serialized)
  213. self.assertEqual('rq_public', response['result']['_meta']['request_id'])
  214. def test_missing_session_on_tool_call_returns_device_protocol_error(self):
  215. gateway = FakeGateway()
  216. handler = PublicMcpHttpHandler(gateway, context_parser=FakeParser())
  217. response = handler.handle_json_rpc(
  218. headers={},
  219. message={
  220. 'jsonrpc': '2.0',
  221. 'id': 1,
  222. 'method': 'tools/call',
  223. 'params': {'name': 'query_order', 'arguments': {'keyword': 'USC'}},
  224. },
  225. )
  226. self.assertEqual(-32001, response['error']['code'])
  227. self.assertIn('Workbuddy', response['error']['message'])
  228. self.assertEqual([], gateway.calls)
  229. def test_non_object_tool_params_fail_safely(self):
  230. gateway = FakeGateway()
  231. handler = PublicMcpHttpHandler(gateway, context_parser=FakeParser())
  232. with self.assertLogs('public_server', level='ERROR') as logs:
  233. response = handler.handle_json_rpc(
  234. headers={'X-Gateway-Session': 'GWS_A'},
  235. message={
  236. 'jsonrpc': '2.0',
  237. 'id': 9,
  238. 'method': 'tools/call',
  239. 'params': 'not-an-object',
  240. },
  241. client_ip='10.0.0.5',
  242. )
  243. self.assertTrue(response['result']['isError'])
  244. self.assertEqual(
  245. '工具返回格式异常',
  246. response['result']['structuredContent']['message'],
  247. )
  248. self.assertEqual([], gateway.calls)
  249. record = logs.records[0]
  250. self.assertEqual(9, record.jsonrpc_id)
  251. self.assertTrue(record.request_id.startswith('rq_http_'))
  252. self.assertEqual('', record.tool_code)
  253. self.assertEqual('MCP_9001', record.response_code)
  254. self.assertEqual('PARAM_VALIDATION_FAILED', record.diagnostic_reason)
  255. self.assertEqual('ValueError', record.exception_class)
  256. self.assertEqual(record.request_id, response['result']['_meta']['request_id'])
  257. def test_extract_client_ip_ignores_spoofable_forwarded_for_header(self):
  258. client_ip = extract_client_ip(
  259. {'X-Forwarded-For': '203.0.113.9'},
  260. ('10.0.0.5', 54321),
  261. )
  262. self.assertEqual('10.0.0.5', client_ip)
  263. class RateLimitTest(unittest.TestCase):
  264. def _make_handler(self, max_requests=2, max_in_flight=2):
  265. gateway = FakeGateway()
  266. limiter = SimpleRateLimiter(
  267. max_requests=max_requests,
  268. window_seconds=60,
  269. max_in_flight=max_in_flight,
  270. )
  271. return PublicMcpHttpHandler(gateway, context_parser=FakeParser(), rate_limiter=limiter)
  272. def _tools_call_msg(self, tool_name='query_order'):
  273. return {
  274. 'jsonrpc': '2.0', 'id': 1,
  275. 'method': 'tools/call',
  276. 'params': {'name': tool_name, 'arguments': {'keyword': 'test'}},
  277. }
  278. # --- tools/call per-tool独立限流 ---
  279. def test_known_tool_rate_limited_after_quota_exhausted(self):
  280. handler = self._make_handler(max_requests=1)
  281. handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg(), client_ip='1.2.3.4')
  282. with self.assertLogs('public_server', level='WARNING') as logs:
  283. response = handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg(), client_ip='1.2.3.4')
  284. self.assertIn('error', response)
  285. self.assertEqual(-32029, response['error']['code'])
  286. self.assertIn('Rate limit', response['error']['message'])
  287. record = logs.records[0]
  288. self.assertEqual(1, record.jsonrpc_id)
  289. self.assertEqual('query_order', record.tool_code)
  290. self.assertEqual(-32029, record.protocol_code)
  291. self.assertEqual('RATE_LIMIT_EXCEEDED', record.diagnostic_reason)
  292. self.assertEqual(response['error']['data']['request_id'], record.request_id)
  293. def test_two_known_tools_have_independent_quotas(self):
  294. # query_order 限流不影响 query_track
  295. handler = self._make_handler(max_requests=1)
  296. handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg('query_order'), client_ip='1.2.3.4')
  297. # query_order quota exhausted, query_track should still work
  298. response = handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg('query_track'), client_ip='1.2.3.4')
  299. self.assertNotIn('error', response)
  300. def test_unknown_tool_name_uses_shared_session_bucket_not_new_bucket(self):
  301. # 未知工具名应归入 session_id 桶,同一会话的不同未知工具共享同一配额,不能无限新建桶
  302. handler = self._make_handler(max_requests=1)
  303. # 先用 session 桶打一次(用未知工具名触发)
  304. handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg('fake_tool_1'), client_ip='1.2.3.4')
  305. # 再用同会话的另一个未知工具名,应该命中同一个 session 桶,被限流
  306. response = handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg('fake_tool_2'), client_ip='1.2.3.4')
  307. self.assertIn('error', response)
  308. self.assertIn('Rate limit', response['error']['message'])
  309. def test_different_sessions_have_independent_quotas(self):
  310. # 不同会话(不同员工)即使来自同一IP,也拥有独立的配额
  311. handler = self._make_handler(max_requests=1)
  312. handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg(), client_ip='1.1.1.1')
  313. # 不同 session,即使同一 IP,quota 独立 → 应该允许
  314. response = handler.handle_json_rpc({'X-Gateway-Session': 'GWS_B'}, self._tools_call_msg(), client_ip='1.1.1.1')
  315. self.assertNotIn('error', response)
  316. def test_same_session_from_different_ips_shares_quota(self):
  317. # 同一会话从不同IP发来(如移动网络切换),仍共享同一会话配额
  318. handler = self._make_handler(max_requests=1)
  319. handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg(), client_ip='1.1.1.1')
  320. # 同 session,不同 IP → 命中同一 session 桶,被限流
  321. response = handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, self._tools_call_msg(), client_ip='2.2.2.2')
  322. self.assertIn('error', response)
  323. self.assertIn('Rate limit', response['error']['message'])
  324. # --- initialize 和 tools/list 不受限流 ---
  325. def test_initialize_is_never_rate_limited(self):
  326. # initialize 是握手协议,无论请求多少次都不应被限流
  327. handler = self._make_handler(max_requests=1)
  328. for i in range(5):
  329. response = handler.handle_json_rpc({}, {'jsonrpc': '2.0', 'id': i, 'method': 'initialize'}, client_ip='1.2.3.4')
  330. self.assertNotIn('error', response, f'initialize should never be rate limited (attempt {i})')
  331. self.assertIn('protocolVersion', response['result'])
  332. def test_tools_list_is_never_rate_limited(self):
  333. handler = self._make_handler(max_requests=1)
  334. headers = {'X-Gateway-Session': 'GWS_A'}
  335. for request_id in range(1, 4):
  336. response = handler.handle_json_rpc(
  337. headers,
  338. {'jsonrpc': '2.0', 'id': request_id, 'method': 'tools/list'},
  339. client_ip='1.2.3.4',
  340. )
  341. self.assertNotIn('error', response)
  342. def test_tools_list_requests_do_not_consume_tools_call_quota(self):
  343. # tools/list 不进入限流器,因此不会消耗 tools/call 配额
  344. handler = self._make_handler(max_requests=1)
  345. handler.handle_json_rpc({'X-Gateway-Session': 'GWS_A'}, {'jsonrpc': '2.0', 'id': 1, 'method': 'tools/list'}, client_ip='1.2.3.4')
  346. # tools/call 应仍可正常执行(使用 session_id:tool_name bucket)
  347. response = handler.handle_json_rpc(
  348. {'X-Gateway-Session': 'GWS_A'},
  349. {'jsonrpc': '2.0', 'id': 2, 'method': 'tools/call',
  350. 'params': {'name': 'query_order', 'arguments': {'keyword': 'test'}}},
  351. client_ip='1.2.3.4',
  352. )
  353. self.assertNotIn('error', response)
  354. def test_completed_tool_call_releases_in_flight_slot_immediately(self):
  355. handler = self._make_handler(max_requests=0, max_in_flight=1)
  356. headers = {'X-Gateway-Session': 'GWS_A'}
  357. for request_id in range(1, 6):
  358. message = self._tools_call_msg('query_order')
  359. message['id'] = request_id
  360. response = handler.handle_json_rpc(
  361. headers,
  362. message,
  363. client_ip='1.2.3.4',
  364. )
  365. self.assertNotIn('error', response)
  366. def test_exhausted_in_flight_slots_return_rate_limit_error(self):
  367. handler = self._make_handler(max_requests=0, max_in_flight=2)
  368. key = 'GWS_A:query_order'
  369. self.assertTrue(handler.rate_limiter.try_acquire(key))
  370. self.assertTrue(handler.rate_limiter.try_acquire(key))
  371. response = handler.handle_json_rpc(
  372. {'X-Gateway-Session': 'GWS_A'},
  373. self._tools_call_msg('query_order'),
  374. client_ip='1.2.3.4',
  375. )
  376. self.assertEqual(-32029, response['error']['code'])
  377. self.assertIn('in progress', response['error']['message'])
  378. def test_failed_tool_call_releases_in_flight_slot(self):
  379. handler = self._make_handler(max_requests=0, max_in_flight=1)
  380. original_call = handler.gateway_app.call_tool
  381. def fail(*args, **kwargs):
  382. raise RuntimeError('backend failed')
  383. handler.gateway_app.call_tool = fail
  384. handler.handle_json_rpc(
  385. {'X-Gateway-Session': 'GWS_A'},
  386. self._tools_call_msg('query_order'),
  387. client_ip='1.2.3.4',
  388. )
  389. handler.gateway_app.call_tool = original_call
  390. response = handler.handle_json_rpc(
  391. {'X-Gateway-Session': 'GWS_A'},
  392. self._tools_call_msg('query_order'),
  393. client_ip='1.2.3.4',
  394. )
  395. self.assertNotIn('error', response)
  396. class HttpDisconnectTest(unittest.TestCase):
  397. def test_broken_pipe_while_writing_response_is_handled(self):
  398. handler_class = create_http_handler(FakeGateway(), rate_limiter=None)
  399. handler = object.__new__(handler_class)
  400. handler.send_response = Mock()
  401. handler.send_header = Mock()
  402. handler.end_headers = Mock()
  403. handler.wfile = Mock()
  404. handler.wfile.write.side_effect = BrokenPipeError()
  405. with self.assertLogs('public_server', level='INFO') as logs:
  406. handler._write_json({'jsonrpc': '2.0', 'id': 1, 'result': {}})
  407. self.assertIn('client disconnected before response', logs.output[0])
  408. if __name__ == '__main__':
  409. unittest.main()