import unittest from unittest.mock import MagicMock from constants import DEVICE_INVALID_MESSAGE from public_gateway import PublicGatewayApp class TestPublicGatewayApp(unittest.TestCase): def setUp(self): self.mock_session_store = MagicMock() self.mock_api_client = MagicMock() self.mock_api_client.list_enabled_tools.return_value = { 'code': 'MCP_0000', 'data': { 'tool_codes': [ 'query_order', 'query_track', 'query_order_exact', 'list_order_filter_options', ], }, } self.app = PublicGatewayApp( session_store=self.mock_session_store, api_client=self.mock_api_client, ) def test_list_tools_returns_metadata(self): self.mock_session_store.get.return_value = { 'mcp_token': 'MT_token', } self.mock_api_client.list_enabled_tools.return_value = { 'code': 'MCP_0000', 'data': { 'tool_codes': [ 'query_order', 'query_track', 'query_order_exact', 'list_order_filter_options', ], }, } tools = self.app.list_tools('GWS_test') self.assertIsInstance(tools, list) 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.assertIn('query_order_exact', tool_names) self.assertIn('list_order_filter_options', tool_names) self.assertNotIn('bind_auth_code', tool_names) for tool in tools: self.assertIn('name', tool) self.assertIn('description', tool) self.assertIn('input_schema', tool) def test_list_tools_intersects_enabled_codes_and_preserves_local_order(self): self.mock_session_store.get.return_value = { 'mcp_token': 'MT_token', } self.mock_api_client.list_enabled_tools.return_value = { 'code': 'MCP_0000', 'data': { 'tool_codes': [ 'query_track', 'query_order_exact', 'unknown_tool', ], }, } tools = self.app.list_tools('GWS_test') self.assertEqual( ['query_track', 'query_order_exact'], [tool['name'] for tool in tools], ) self.mock_api_client.list_enabled_tools.assert_called_once_with('MT_token') def test_list_tools_rejects_missing_session_without_querying_registry(self): self.mock_session_store.get.return_value = None with self.assertRaisesRegex(RuntimeError, DEVICE_INVALID_MESSAGE): self.app.list_tools('GWS_missing') self.mock_api_client.list_enabled_tools.assert_not_called() def test_list_tools_fails_closed_on_registry_error(self): self.mock_session_store.get.return_value = {'mcp_token': 'MT_token'} self.mock_api_client.list_enabled_tools.return_value = { 'code': 'MCP_9001', 'msg': 'registry unavailable', 'data': {}, } with self.assertRaisesRegex(RuntimeError, 'registry unavailable'): self.app.list_tools('GWS_test') def test_enabled_tool_names_rejects_non_dict_response(self): with self.assertRaisesRegex(RuntimeError, 'invalid enabled tool response'): self.app._enabled_tool_names(None) def test_enabled_tool_names_rejects_missing_tool_code_list(self): response = {'code': 'MCP_0000', 'data': {}} with self.assertRaisesRegex(RuntimeError, 'invalid enabled tool response'): self.app._enabled_tool_names(response) def test_build_request_id_generates_id_when_empty(self): request_id = self.app.build_request_id('') self.assertTrue(request_id.startswith('rq_')) self.assertEqual(len(request_id), 3 + 16) def test_build_request_id_uses_provided_id(self): provided_id = 'custom_request_123' request_id = self.app.build_request_id(provided_id) self.assertEqual(request_id, provided_id) def test_call_tool_bind_auth_code_is_not_registered_in_public_mode(self): with self.assertRaises(KeyError) as context: self.app.call_tool('GWS_bind', 'bind_auth_code', {'auth_code': 'AC_test'}) self.assertIn('tool not registered', str(context.exception)) def test_call_tool_raises_on_unregistered_tool(self): with self.assertRaises(KeyError) as context: self.app.call_tool('GWS_test', 'nonexistent_tool', {}) self.assertIn('tool not registered', str(context.exception)) def test_call_tool_raises_when_session_not_found(self): self.mock_session_store.get.return_value = None with self.assertRaises(RuntimeError) as context: self.app.call_tool('GWS_nosession', 'query_order', {}) self.assertEqual(str(context.exception), DEVICE_INVALID_MESSAGE) def test_call_tool_raises_when_no_token_in_session(self): self.mock_session_store.get.return_value = { 'admin_id': 123, 'company_id': 100 } with self.assertRaises(RuntimeError) as context: self.app.call_tool('GWS_notoken', 'query_order', {}) self.assertEqual(str(context.exception), DEVICE_INVALID_MESSAGE) def test_call_tool_success(self): gateway_session_id = 'GWS_valid' tool_name = 'query_order' arguments = {'order_no': 'ABC123'} self.mock_session_store.get.return_value = { 'mcp_token': 'MT_token123', 'admin_id': 456, 'company_id': 200 } self.mock_api_client.call_tool.return_value = { 'code': '0', 'msg': 'success', 'data': {'order': 'details'} } self.mock_api_client.list_enabled_tools.return_value = { 'code': 'MCP_0000', 'data': {'tool_codes': ['query_order']}, } result = self.app.call_tool(gateway_session_id, tool_name, arguments, 'rq_test') self.mock_api_client.call_tool.assert_called_once() call_args = self.mock_api_client.call_tool.call_args[1] self.assertEqual(call_args['token'], 'MT_token123') self.assertEqual(call_args['tool_code'], 'query_order') self.assertEqual(call_args['payload'], arguments) self.assertEqual(call_args['request_id'], 'rq_test') self.assertEqual(result['code'], '0') def test_call_tool_succeeds_when_store_has_no_touch_method(self): class ReadOnlySessionStore: def get(self, gateway_session_id): return { 'mcp_token': 'MT_read_only', 'admin_id': 1, 'company_id': 2, } api_client = MagicMock() api_client.list_enabled_tools.return_value = { 'code': 'MCP_0000', 'data': {'tool_codes': ['query_order']}, } api_client.call_tool.return_value = {'code': 'MCP_0000'} app = PublicGatewayApp(ReadOnlySessionStore(), api_client) result = app.call_tool( 'GWS_read_only', 'query_order', {'keyword': 'ORDER-1'}, 'rq_read_only', ) self.assertEqual('MCP_0000', result['code']) api_client.call_tool.assert_called_once() def test_call_tool_rejects_dynamically_disabled_tool_before_forwarding(self): self.mock_session_store.get.return_value = { 'mcp_token': 'MT_token', 'admin_id': 1, 'company_id': 1, } self.mock_api_client.list_enabled_tools.return_value = { 'code': 'MCP_0000', 'data': {'tool_codes': ['query_track']}, } with self.assertRaisesRegex(RuntimeError, 'tool disabled: query_order'): self.app.call_tool('GWS_test', 'query_order', {}) self.mock_session_store.get.assert_called_once_with('GWS_test') self.mock_api_client.call_tool.assert_not_called() def test_call_tool_generates_request_id_when_not_provided(self): 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'} self.app.call_tool('GWS_test', 'query_order', {}, '') call_args = self.mock_api_client.call_tool.call_args[1] self.assertTrue(call_args['request_id'].startswith('rq_')) def test_call_tool_uses_provided_request_id(self): 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'} custom_request_id = 'rq_custom_123' self.app.call_tool('GWS_test', 'query_order', {}, custom_request_id) call_args = self.mock_api_client.call_tool.call_args[1] self.assertEqual(call_args['request_id'], custom_request_id) def test_call_tool_forwards_client_ip_to_api_client(self): 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'} self.app.call_tool('GWS_test', 'query_order', {}, 'rq_client_ip', client_ip='203.0.113.9') call_args = self.mock_api_client.call_tool.call_args[1] self.assertEqual(call_args['client_ip'], '203.0.113.9') def test_call_tool_logs_error_on_exception(self): gateway_session_id = 'GWS_error_test' tool_name = 'query_order' self.mock_session_store.get.return_value = { 'mcp_token': 'MT_token', 'admin_id': 999, 'company_id': 888 } self.mock_api_client.call_tool.side_effect = RuntimeError('API connection failed') with self.assertRaises(RuntimeError) as context: self.app.call_tool(gateway_session_id, tool_name, {}, 'rq_err') self.assertIn('API connection failed', str(context.exception)) if __name__ == '__main__': unittest.main()