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()