test_public_server_integration.py 8.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255
  1. import json
  2. import unittest
  3. from http.client import HTTPConnection
  4. from threading import Thread
  5. from time import sleep
  6. from unittest.mock import MagicMock
  7. from constants import DEVICE_INVALID_MESSAGE
  8. from public_gateway import PublicGatewayApp
  9. from public_server import serve_public
  10. from utils.security import generate_gateway_session_id
  11. class TestPublicServerIntegration(unittest.TestCase):
  12. @classmethod
  13. def setUpClass(cls):
  14. # Create mock dependencies
  15. cls.mock_session_store = MagicMock()
  16. cls.mock_api_client = MagicMock()
  17. cls.mock_auth_client = MagicMock()
  18. # Create gateway app
  19. cls.gateway_app = PublicGatewayApp(
  20. session_store=cls.mock_session_store,
  21. api_client=cls.mock_api_client,
  22. auth_client=cls.mock_auth_client
  23. )
  24. # Start server in background thread
  25. cls.server_thread = Thread(
  26. target=serve_public,
  27. args=(cls.gateway_app,),
  28. kwargs={'host': '127.0.0.1', 'port': 18765, 'enable_rate_limit': False},
  29. daemon=True
  30. )
  31. cls.server_thread.start()
  32. # Wait for server to start
  33. sleep(0.5)
  34. def setUp(self):
  35. # Reset mocks before each test
  36. self.mock_session_store.reset_mock()
  37. self.mock_api_client.reset_mock()
  38. self.mock_auth_client.reset_mock()
  39. self.mock_session_store.get.return_value = None
  40. self.mock_api_client.list_enabled_tools.return_value = {
  41. 'code': 'MCP_0000',
  42. 'data': {
  43. 'tool_codes': [
  44. 'query_order',
  45. 'query_track',
  46. 'query_order_exact',
  47. 'list_order_filter_options',
  48. ],
  49. },
  50. }
  51. def _make_request(self, method, params=None, headers=None):
  52. """Helper to make JSON-RPC requests"""
  53. conn = HTTPConnection('127.0.0.1', 18765, timeout=5)
  54. try:
  55. body = json.dumps({
  56. 'jsonrpc': '2.0',
  57. 'id': 1,
  58. 'method': method,
  59. 'params': params or {}
  60. })
  61. request_headers = {'Content-Type': 'application/json'}
  62. if headers:
  63. request_headers.update(headers)
  64. conn.request('POST', '/mcp', body.encode('utf-8'), request_headers)
  65. response = conn.getresponse()
  66. data = response.read().decode('utf-8')
  67. return json.loads(data)
  68. finally:
  69. conn.close()
  70. def test_health_check(self):
  71. conn = HTTPConnection('127.0.0.1', 18765, timeout=5)
  72. try:
  73. conn.request('GET', '/health')
  74. response = conn.getresponse()
  75. data = json.loads(response.read().decode('utf-8'))
  76. self.assertEqual(response.status, 200)
  77. self.assertTrue(data['ok'])
  78. finally:
  79. conn.close()
  80. def test_initialize(self):
  81. result = self._make_request('initialize')
  82. self.assertIn('result', result)
  83. self.assertIn('protocolVersion', result['result'])
  84. self.assertIn('serverInfo', result['result'])
  85. def test_tools_list(self):
  86. session_id = generate_gateway_session_id()
  87. self.mock_session_store.get.return_value = {
  88. 'mcp_token': 'MT_valid_token',
  89. }
  90. result = self._make_request(
  91. 'tools/list',
  92. headers={'X-Gateway-Session': session_id},
  93. )
  94. self.assertIn('result', result)
  95. self.assertIn('tools', result['result'])
  96. tools = result['result']['tools']
  97. self.assertGreater(len(tools), 0)
  98. tool_names = [tool['name'] for tool in tools]
  99. self.assertIn('query_order', tool_names)
  100. self.assertIn('query_track', tool_names)
  101. self.assertIn('query_order_exact', tool_names)
  102. self.assertIn('list_order_filter_options', tool_names)
  103. self.assertNotIn('bind_auth_code', tool_names)
  104. # Check tool structure
  105. for tool in tools:
  106. self.assertIn('name', tool)
  107. self.assertIn('description', tool)
  108. self.assertIn('inputSchema', tool)
  109. by_name = {tool['name']: tool for tool in tools}
  110. exact_tool = by_name['query_order_exact']
  111. self.assertIn('单号类型不明确', exact_tool['description'])
  112. self.assertIn(
  113. '系统订单号',
  114. exact_tool['inputSchema']['properties']['order_number'][
  115. 'description'
  116. ],
  117. )
  118. def test_tools_call_without_session(self):
  119. result = self._make_request('tools/call', {'name': 'query_order', 'arguments': {}})
  120. self.assertIn('result', result)
  121. self.assertTrue(result['result']['isError'])
  122. self.assertIn(DEVICE_INVALID_MESSAGE, result['result']['content'][0]['text'])
  123. def test_tools_call_bind_auth_code_is_not_registered_in_public_mode(self):
  124. session_id = generate_gateway_session_id()
  125. auth_code = 'AC_test123'
  126. result = self._make_request(
  127. 'tools/call',
  128. {
  129. 'name': 'bind_auth_code',
  130. 'arguments': {'auth_code': auth_code}
  131. },
  132. {'X-Gateway-Session': session_id}
  133. )
  134. self.assertIn('result', result)
  135. self.assertTrue(result['result']['isError'])
  136. self.assertNotIn('bind_auth_code', result['result']['content'][0]['text'])
  137. self.assertEqual(
  138. 'MCP_9001',
  139. result['result']['structuredContent']['code'],
  140. )
  141. self.mock_auth_client.exchange.assert_not_called()
  142. def test_tools_call_with_valid_session(self):
  143. session_id = generate_gateway_session_id()
  144. self.mock_session_store.get.return_value = {
  145. 'mcp_token': 'MT_valid_token',
  146. 'admin_id': 456,
  147. 'company_id': 200
  148. }
  149. self.mock_api_client.call_tool.return_value = {
  150. 'code': '0',
  151. 'msg': 'success',
  152. 'data': {'order_no': 'ABC123'}
  153. }
  154. result = self._make_request(
  155. 'tools/call',
  156. {
  157. 'name': 'query_order',
  158. 'arguments': {'order_no': 'ABC123'}
  159. },
  160. {'X-Gateway-Session': session_id}
  161. )
  162. self.assertIn('result', result)
  163. self.assertFalse(result['result']['isError'])
  164. self.mock_api_client.call_tool.assert_called_once()
  165. def test_method_not_found(self):
  166. result = self._make_request('nonexistent_method')
  167. self.assertIn('error', result)
  168. self.assertEqual(result['error']['code'], -32601)
  169. self.assertIn('Method not found', result['error']['message'])
  170. def test_404_for_invalid_path(self):
  171. conn = HTTPConnection('127.0.0.1', 18765, timeout=5)
  172. try:
  173. conn.request('POST', '/invalid')
  174. response = conn.getresponse()
  175. self.assertEqual(response.status, 404)
  176. finally:
  177. conn.close()
  178. def test_session_from_authorization_header(self):
  179. session_id = generate_gateway_session_id()
  180. self.mock_session_store.get.return_value = {
  181. 'mcp_token': 'MT_token',
  182. 'admin_id': 1,
  183. 'company_id': 1
  184. }
  185. self.mock_api_client.call_tool.return_value = {'code': '0', 'data': {}}
  186. result = self._make_request(
  187. 'tools/call',
  188. {'name': 'query_order', 'arguments': {}},
  189. {'Authorization': f'Bearer {session_id}'}
  190. )
  191. self.assertIn('result', result)
  192. self.assertFalse(result['result']['isError'])
  193. def test_session_from_cookie(self):
  194. session_id = generate_gateway_session_id()
  195. self.mock_session_store.get.return_value = {
  196. 'mcp_token': 'MT_token',
  197. 'admin_id': 1,
  198. 'company_id': 1
  199. }
  200. self.mock_api_client.call_tool.return_value = {'code': '0', 'data': {}}
  201. result = self._make_request(
  202. 'tools/call',
  203. {'name': 'query_order', 'arguments': {}},
  204. {'Cookie': f'gateway_session_id={session_id}; other=value'}
  205. )
  206. self.assertIn('result', result)
  207. self.assertFalse(result['result']['isError'])
  208. if __name__ == '__main__':
  209. unittest.main()