| 12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970717273747576777879 |
- import unittest
- from public_server import PublicMcpHttpHandler, extract_client_ip
- class FakeContext:
- def __init__(self, gateway_session_id):
- self.gateway_session_id = gateway_session_id
- def has_session(self):
- return bool(self.gateway_session_id)
- class FakeParser:
- def parse(self, headers):
- return FakeContext(headers.get('X-Gateway-Session', ''))
- class FakeGateway:
- def __init__(self):
- self.calls = []
- def list_tools(self):
- return [{'name': 'query_order', 'description': 'query order', 'input_schema': {'type': 'object'}}]
- 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}}
- class PublicMcpHttpHandlerTest(unittest.TestCase):
- 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'},
- message={
- 'jsonrpc': '2.0',
- 'id': 1,
- 'method': 'tools/call',
- 'params': {
- 'name': 'query_order',
- 'arguments': {'keyword': 'USC'},
- },
- },
- )
- self.assertEqual(False, response['result']['isError'])
- self.assertEqual('GWS_A', gateway.calls[0][0])
- self.assertEqual('query_order', gateway.calls[0][1])
- def test_missing_session_on_tool_call_returns_error_content(self):
- gateway = FakeGateway()
- handler = PublicMcpHttpHandler(gateway, context_parser=FakeParser())
- response = handler.handle_json_rpc(
- headers={},
- message={
- 'jsonrpc': '2.0',
- 'id': 1,
- 'method': 'tools/call',
- 'params': {'name': 'query_order', 'arguments': {'keyword': 'USC'}},
- },
- )
- self.assertTrue(response['result']['isError'])
- self.assertIn('绑定授权码', response['result']['content'][0]['text'])
- def test_extract_client_ip_ignores_spoofable_forwarded_for_header(self):
- client_ip = extract_client_ip(
- {'X-Forwarded-For': '203.0.113.9'},
- ('10.0.0.5', 54321),
- )
- self.assertEqual('10.0.0.5', client_ip)
- if __name__ == '__main__':
- unittest.main()
|