test_public_server_integration.py 6.9 KB

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