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.app = PublicGatewayApp( session_store=self.mock_session_store, api_client=self.mock_api_client, ) def test_list_tools_returns_metadata(self): tools = self.app.list_tools() 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.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_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'} } 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_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_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()