test_public_gateway_unit.py 5.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151
  1. import unittest
  2. from unittest.mock import MagicMock
  3. from constants import DEVICE_INVALID_MESSAGE
  4. from public_gateway import PublicGatewayApp
  5. class TestPublicGatewayApp(unittest.TestCase):
  6. def setUp(self):
  7. self.mock_session_store = MagicMock()
  8. self.mock_api_client = MagicMock()
  9. self.app = PublicGatewayApp(
  10. session_store=self.mock_session_store,
  11. api_client=self.mock_api_client,
  12. )
  13. def test_list_tools_returns_metadata(self):
  14. tools = self.app.list_tools()
  15. self.assertIsInstance(tools, list)
  16. self.assertGreater(len(tools), 0)
  17. tool_names = [tool['name'] for tool in tools]
  18. self.assertIn('query_order', tool_names)
  19. self.assertIn('query_track', tool_names)
  20. self.assertNotIn('bind_auth_code', tool_names)
  21. for tool in tools:
  22. self.assertIn('name', tool)
  23. self.assertIn('description', tool)
  24. self.assertIn('input_schema', tool)
  25. def test_build_request_id_generates_id_when_empty(self):
  26. request_id = self.app.build_request_id('')
  27. self.assertTrue(request_id.startswith('rq_'))
  28. self.assertEqual(len(request_id), 3 + 16)
  29. def test_build_request_id_uses_provided_id(self):
  30. provided_id = 'custom_request_123'
  31. request_id = self.app.build_request_id(provided_id)
  32. self.assertEqual(request_id, provided_id)
  33. def test_call_tool_bind_auth_code_is_not_registered_in_public_mode(self):
  34. with self.assertRaises(KeyError) as context:
  35. self.app.call_tool('GWS_bind', 'bind_auth_code', {'auth_code': 'AC_test'})
  36. self.assertIn('tool not registered', str(context.exception))
  37. def test_call_tool_raises_on_unregistered_tool(self):
  38. with self.assertRaises(KeyError) as context:
  39. self.app.call_tool('GWS_test', 'nonexistent_tool', {})
  40. self.assertIn('tool not registered', str(context.exception))
  41. def test_call_tool_raises_when_session_not_found(self):
  42. self.mock_session_store.get.return_value = None
  43. with self.assertRaises(RuntimeError) as context:
  44. self.app.call_tool('GWS_nosession', 'query_order', {})
  45. self.assertEqual(str(context.exception), DEVICE_INVALID_MESSAGE)
  46. def test_call_tool_raises_when_no_token_in_session(self):
  47. self.mock_session_store.get.return_value = {
  48. 'admin_id': 123,
  49. 'company_id': 100
  50. }
  51. with self.assertRaises(RuntimeError) as context:
  52. self.app.call_tool('GWS_notoken', 'query_order', {})
  53. self.assertEqual(str(context.exception), DEVICE_INVALID_MESSAGE)
  54. def test_call_tool_success(self):
  55. gateway_session_id = 'GWS_valid'
  56. tool_name = 'query_order'
  57. arguments = {'order_no': 'ABC123'}
  58. self.mock_session_store.get.return_value = {
  59. 'mcp_token': 'MT_token123',
  60. 'admin_id': 456,
  61. 'company_id': 200
  62. }
  63. self.mock_api_client.call_tool.return_value = {
  64. 'code': '0',
  65. 'msg': 'success',
  66. 'data': {'order': 'details'}
  67. }
  68. result = self.app.call_tool(gateway_session_id, tool_name, arguments, 'rq_test')
  69. self.mock_api_client.call_tool.assert_called_once()
  70. call_args = self.mock_api_client.call_tool.call_args[1]
  71. self.assertEqual(call_args['token'], 'MT_token123')
  72. self.assertEqual(call_args['tool_code'], 'query_order')
  73. self.assertEqual(call_args['payload'], arguments)
  74. self.assertEqual(call_args['request_id'], 'rq_test')
  75. self.assertEqual(result['code'], '0')
  76. def test_call_tool_generates_request_id_when_not_provided(self):
  77. self.mock_session_store.get.return_value = {
  78. 'mcp_token': 'MT_token',
  79. 'admin_id': 1,
  80. 'company_id': 1
  81. }
  82. self.mock_api_client.call_tool.return_value = {'code': '0'}
  83. self.app.call_tool('GWS_test', 'query_order', {}, '')
  84. call_args = self.mock_api_client.call_tool.call_args[1]
  85. self.assertTrue(call_args['request_id'].startswith('rq_'))
  86. def test_call_tool_uses_provided_request_id(self):
  87. self.mock_session_store.get.return_value = {
  88. 'mcp_token': 'MT_token',
  89. 'admin_id': 1,
  90. 'company_id': 1
  91. }
  92. self.mock_api_client.call_tool.return_value = {'code': '0'}
  93. custom_request_id = 'rq_custom_123'
  94. self.app.call_tool('GWS_test', 'query_order', {}, custom_request_id)
  95. call_args = self.mock_api_client.call_tool.call_args[1]
  96. self.assertEqual(call_args['request_id'], custom_request_id)
  97. def test_call_tool_logs_error_on_exception(self):
  98. gateway_session_id = 'GWS_error_test'
  99. tool_name = 'query_order'
  100. self.mock_session_store.get.return_value = {
  101. 'mcp_token': 'MT_token',
  102. 'admin_id': 999,
  103. 'company_id': 888
  104. }
  105. self.mock_api_client.call_tool.side_effect = RuntimeError('API connection failed')
  106. with self.assertRaises(RuntimeError) as context:
  107. self.app.call_tool(gateway_session_id, tool_name, {}, 'rq_err')
  108. self.assertIn('API connection failed', str(context.exception))
  109. if __name__ == '__main__':
  110. unittest.main()