test_public_server.py 2.5 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970717273747576777879
  1. import unittest
  2. from public_server import PublicMcpHttpHandler, extract_client_ip
  3. class FakeContext:
  4. def __init__(self, gateway_session_id):
  5. self.gateway_session_id = gateway_session_id
  6. def has_session(self):
  7. return bool(self.gateway_session_id)
  8. class FakeParser:
  9. def parse(self, headers):
  10. return FakeContext(headers.get('X-Gateway-Session', ''))
  11. class FakeGateway:
  12. def __init__(self):
  13. self.calls = []
  14. def list_tools(self):
  15. return [{'name': 'query_order', 'description': 'query order', 'input_schema': {'type': 'object'}}]
  16. def call_tool(self, gateway_session_id, name, arguments=None, request_id=''):
  17. self.calls.append((gateway_session_id, name, arguments, request_id))
  18. return {'code': 'MCP_0000', 'data': {'ok': True}}
  19. class PublicMcpHttpHandlerTest(unittest.TestCase):
  20. def test_handle_tools_call_passes_gateway_session_to_public_gateway(self):
  21. gateway = FakeGateway()
  22. handler = PublicMcpHttpHandler(gateway, context_parser=FakeParser())
  23. response = handler.handle_json_rpc(
  24. headers={'X-Gateway-Session': 'GWS_A'},
  25. message={
  26. 'jsonrpc': '2.0',
  27. 'id': 1,
  28. 'method': 'tools/call',
  29. 'params': {
  30. 'name': 'query_order',
  31. 'arguments': {'keyword': 'USC'},
  32. },
  33. },
  34. )
  35. self.assertEqual(False, response['result']['isError'])
  36. self.assertEqual('GWS_A', gateway.calls[0][0])
  37. self.assertEqual('query_order', gateway.calls[0][1])
  38. def test_missing_session_on_tool_call_returns_error_content(self):
  39. gateway = FakeGateway()
  40. handler = PublicMcpHttpHandler(gateway, context_parser=FakeParser())
  41. response = handler.handle_json_rpc(
  42. headers={},
  43. message={
  44. 'jsonrpc': '2.0',
  45. 'id': 1,
  46. 'method': 'tools/call',
  47. 'params': {'name': 'query_order', 'arguments': {'keyword': 'USC'}},
  48. },
  49. )
  50. self.assertTrue(response['result']['isError'])
  51. self.assertIn('这台设备的 Workbuddy 配置已失效,请重新生成配置', response['result']['content'][0]['text'])
  52. def test_extract_client_ip_ignores_spoofable_forwarded_for_header(self):
  53. client_ip = extract_client_ip(
  54. {'X-Forwarded-For': '203.0.113.9'},
  55. ('10.0.0.5', 54321),
  56. )
  57. self.assertEqual('10.0.0.5', client_ip)
  58. if __name__ == '__main__':
  59. unittest.main()