| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219 |
- import json
- import unittest
- from http.client import HTTPConnection
- from threading import Thread
- from time import sleep
- from unittest.mock import MagicMock
- from constants import DEVICE_INVALID_MESSAGE
- from public_gateway import PublicGatewayApp
- from public_server import serve_public
- from utils.security import generate_gateway_session_id
- class TestPublicServerIntegration(unittest.TestCase):
- @classmethod
- def setUpClass(cls):
- # Create mock dependencies
- cls.mock_session_store = MagicMock()
- cls.mock_api_client = MagicMock()
- cls.mock_auth_client = MagicMock()
- # Create gateway app
- cls.gateway_app = PublicGatewayApp(
- session_store=cls.mock_session_store,
- api_client=cls.mock_api_client,
- auth_client=cls.mock_auth_client
- )
- # Start server in background thread
- cls.server_thread = Thread(
- target=serve_public,
- args=(cls.gateway_app,),
- kwargs={'host': '127.0.0.1', 'port': 18765, 'enable_rate_limit': False},
- daemon=True
- )
- cls.server_thread.start()
- # Wait for server to start
- sleep(0.5)
- def setUp(self):
- # Reset mocks before each test
- self.mock_session_store.reset_mock()
- self.mock_api_client.reset_mock()
- self.mock_auth_client.reset_mock()
- def _make_request(self, method, params=None, headers=None):
- """Helper to make JSON-RPC requests"""
- conn = HTTPConnection('127.0.0.1', 18765, timeout=5)
- try:
- body = json.dumps({
- 'jsonrpc': '2.0',
- 'id': 1,
- 'method': method,
- 'params': params or {}
- })
- request_headers = {'Content-Type': 'application/json'}
- if headers:
- request_headers.update(headers)
- conn.request('POST', '/mcp', body.encode('utf-8'), request_headers)
- response = conn.getresponse()
- data = response.read().decode('utf-8')
- return json.loads(data)
- finally:
- conn.close()
- def test_health_check(self):
- conn = HTTPConnection('127.0.0.1', 18765, timeout=5)
- try:
- conn.request('GET', '/health')
- response = conn.getresponse()
- data = json.loads(response.read().decode('utf-8'))
- self.assertEqual(response.status, 200)
- self.assertTrue(data['ok'])
- finally:
- conn.close()
- def test_initialize(self):
- result = self._make_request('initialize')
- self.assertIn('result', result)
- self.assertIn('protocolVersion', result['result'])
- self.assertIn('serverInfo', result['result'])
- def test_tools_list(self):
- result = self._make_request('tools/list')
- self.assertIn('result', result)
- self.assertIn('tools', result['result'])
- tools = result['result']['tools']
- self.assertGreater(len(tools), 0)
- tool_names = [tool['name'] for tool in tools]
- self.assertIn('query_order', tool_names)
- self.assertIn('query_track', tool_names)
- self.assertNotIn('bind_auth_code', tool_names)
- # Check tool structure
- for tool in tools:
- self.assertIn('name', tool)
- self.assertIn('description', tool)
- self.assertIn('inputSchema', tool)
- def test_tools_call_without_session(self):
- result = self._make_request('tools/call', {'name': 'query_order', 'arguments': {}})
- self.assertIn('result', result)
- self.assertTrue(result['result']['isError'])
- self.assertIn(DEVICE_INVALID_MESSAGE, result['result']['content'][0]['text'])
- def test_tools_call_bind_auth_code_is_not_registered_in_public_mode(self):
- session_id = generate_gateway_session_id()
- auth_code = 'AC_test123'
- result = self._make_request(
- 'tools/call',
- {
- 'name': 'bind_auth_code',
- 'arguments': {'auth_code': auth_code}
- },
- {'X-Gateway-Session': session_id}
- )
- self.assertIn('result', result)
- self.assertTrue(result['result']['isError'])
- self.assertIn('tool not registered', result['result']['content'][0]['text'])
- self.mock_auth_client.exchange.assert_not_called()
- def test_tools_call_with_valid_session(self):
- session_id = generate_gateway_session_id()
- self.mock_session_store.get.return_value = {
- 'mcp_token': 'MT_valid_token',
- 'admin_id': 456,
- 'company_id': 200
- }
- self.mock_api_client.call_tool.return_value = {
- 'code': '0',
- 'msg': 'success',
- 'data': {'order_no': 'ABC123'}
- }
- result = self._make_request(
- 'tools/call',
- {
- 'name': 'query_order',
- 'arguments': {'order_no': 'ABC123'}
- },
- {'X-Gateway-Session': session_id}
- )
- self.assertIn('result', result)
- self.assertFalse(result['result']['isError'])
- self.mock_api_client.call_tool.assert_called_once()
- def test_method_not_found(self):
- result = self._make_request('nonexistent_method')
- self.assertIn('error', result)
- self.assertEqual(result['error']['code'], -32601)
- self.assertIn('Method not found', result['error']['message'])
- def test_404_for_invalid_path(self):
- conn = HTTPConnection('127.0.0.1', 18765, timeout=5)
- try:
- conn.request('POST', '/invalid')
- response = conn.getresponse()
- self.assertEqual(response.status, 404)
- finally:
- conn.close()
- def test_session_from_authorization_header(self):
- session_id = generate_gateway_session_id()
- self.mock_session_store.get.return_value = {
- 'mcp_token': 'MT_token',
- 'admin_id': 1,
- 'company_id': 1
- }
- self.mock_api_client.call_tool.return_value = {'code': '0', 'data': {}}
- result = self._make_request(
- 'tools/call',
- {'name': 'query_order', 'arguments': {}},
- {'Authorization': f'Bearer {session_id}'}
- )
- self.assertIn('result', result)
- self.assertFalse(result['result']['isError'])
- def test_session_from_cookie(self):
- session_id = generate_gateway_session_id()
- self.mock_session_store.get.return_value = {
- 'mcp_token': 'MT_token',
- 'admin_id': 1,
- 'company_id': 1
- }
- self.mock_api_client.call_tool.return_value = {'code': '0', 'data': {}}
- result = self._make_request(
- 'tools/call',
- {'name': 'query_order', 'arguments': {}},
- {'Cookie': f'gateway_session_id={session_id}; other=value'}
- )
- self.assertIn('result', result)
- self.assertFalse(result['result']['isError'])
- if __name__ == '__main__':
- unittest.main()
|