test_public_gateway.py 10 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280
  1. import unittest
  2. from public_gateway import PublicGatewayApp
  3. from services.diagnostic_event import RequestDiagnosticEmitter
  4. class FakeSessionStore:
  5. def __init__(self):
  6. self.sessions = {}
  7. self.touched = []
  8. def get(self, gateway_session_id):
  9. return self.sessions.get(gateway_session_id)
  10. def touch_session(self, gateway_session_id):
  11. self.touched.append(gateway_session_id)
  12. class FakeApiClient:
  13. def __init__(self):
  14. self.calls = []
  15. def call_tool(self, token, tool_code, route_path, payload, request_id, client_ip=''):
  16. self.calls.append((token, tool_code, route_path, payload, request_id, client_ip))
  17. return {'code': 'MCP_0000', 'data': {'token_used': token}}
  18. def list_enabled_tools(self, token, request_id=''):
  19. return {
  20. 'code': 'MCP_0000',
  21. 'data': {
  22. 'tool_codes': [
  23. 'query_order',
  24. 'query_track',
  25. 'query_order_exact',
  26. 'list_order_filter_options',
  27. ],
  28. },
  29. }
  30. class PublicGatewayAppTest(unittest.TestCase):
  31. def test_missing_redis_session_emits_gateway_session_failure(self):
  32. reporter = RecordingReporter()
  33. emitter = RequestDiagnosticEmitter(
  34. reporter,
  35. 'rq_missing',
  36. defer_until_identity=True,
  37. )
  38. app = PublicGatewayApp(FakeSessionStore(), FakeApiClient())
  39. with self.assertRaises(RuntimeError):
  40. app.call_tool(
  41. 'GWS_missing',
  42. 'query_order',
  43. request_id='rq_missing',
  44. diagnostic_emitter=emitter,
  45. )
  46. emitter.flush()
  47. event = reporter.events[-1]
  48. self.assertEqual('gateway_session', event['stage'])
  49. self.assertEqual('failed', event['status'])
  50. self.assertEqual('GATEWAY_SESSION_NOT_FOUND', event['event_code'])
  51. self.assertIn('session_hash', event)
  52. def test_enabled_tool_lookup_failure_emits_backend_failure(self):
  53. class FailingListClient(FakeApiClient):
  54. def list_enabled_tools(self, token, request_id=''):
  55. raise OSError('registry unavailable secret')
  56. store = FakeSessionStore()
  57. store.sessions['GWS_A'] = {
  58. 'mcp_token': 'MT_A',
  59. 'admin_id': 88,
  60. 'company_id': 1002,
  61. }
  62. reporter = RecordingReporter()
  63. emitter = RequestDiagnosticEmitter(
  64. reporter,
  65. 'rq_lookup',
  66. defer_until_identity=True,
  67. )
  68. emitter.emit(
  69. stage='request_ingress',
  70. status='started',
  71. event_code='REQUEST_RECEIVED',
  72. context={'transport': 'http'},
  73. )
  74. app = PublicGatewayApp(store, FailingListClient())
  75. with self.assertRaises(OSError):
  76. app.call_tool(
  77. 'GWS_A',
  78. 'query_order',
  79. request_id='rq_lookup',
  80. diagnostic_emitter=emitter,
  81. )
  82. self.assertTrue(all(
  83. event['company_id'] == 1002 for event in reporter.events
  84. ))
  85. failed = reporter.events[-1]
  86. self.assertEqual('backend_call', failed['stage'])
  87. self.assertEqual('failed', failed['status'])
  88. self.assertEqual('ENABLED_TOOL_LOOKUP_FAILED', failed['event_code'])
  89. def test_unknown_and_disabled_tools_emit_backend_failures(self):
  90. store = FakeSessionStore()
  91. store.sessions['GWS_A'] = {
  92. 'mcp_token': 'MT_A',
  93. 'admin_id': 88,
  94. 'company_id': 1002,
  95. }
  96. app = PublicGatewayApp(store, FakeApiClient())
  97. cases = (
  98. ('not_registered', 'TOOL_NOT_REGISTERED', KeyError),
  99. ('query_outbound_detail', 'TOOL_DISABLED', RuntimeError),
  100. )
  101. for tool_name, event_code, exception_class in cases:
  102. with self.subTest(tool_name=tool_name):
  103. reporter = RecordingReporter()
  104. emitter = RequestDiagnosticEmitter(
  105. reporter,
  106. 'rq_tool_check',
  107. defer_until_identity=True,
  108. )
  109. with self.assertRaises(exception_class):
  110. app.call_tool(
  111. 'GWS_A',
  112. tool_name,
  113. request_id='rq_tool_check',
  114. diagnostic_emitter=emitter,
  115. )
  116. self.assertEqual(event_code, reporter.events[-1]['event_code'])
  117. self.assertEqual('failed', reporter.events[-1]['status'])
  118. def test_enabled_tool_lookup_failure_without_emitter_keeps_exception(self):
  119. class FailingListClient(FakeApiClient):
  120. def list_enabled_tools(self, token, request_id=''):
  121. raise OSError('registry unavailable')
  122. store = FakeSessionStore()
  123. store.sessions['GWS_A'] = {'mcp_token': 'MT_A'}
  124. app = PublicGatewayApp(store, FailingListClient())
  125. with self.assertRaises(OSError):
  126. app.call_tool('GWS_A', 'query_order', request_id='rq_lookup')
  127. def test_diagnostic_events_include_session_identity_and_backend_result(self):
  128. store = FakeSessionStore()
  129. store.sessions['GWS_A'] = {
  130. 'mcp_token': 'MT_A',
  131. 'admin_id': 88,
  132. 'company_id': 1002,
  133. }
  134. reporter = RecordingReporter()
  135. emitter = RequestDiagnosticEmitter(reporter, 'rq_a')
  136. app = PublicGatewayApp(
  137. session_store=store,
  138. api_client=FakeApiClient(),
  139. auth_client=None,
  140. )
  141. app.call_tool(
  142. 'GWS_A',
  143. 'query_order',
  144. {'keyword': 'A'},
  145. request_id='rq_a',
  146. diagnostic_emitter=emitter,
  147. )
  148. self.assertEqual(
  149. ['gateway_session', 'backend_call', 'backend_call'],
  150. [event['stage'] for event in reporter.events],
  151. )
  152. self.assertEqual(
  153. ['succeeded', 'started', 'succeeded'],
  154. [event['status'] for event in reporter.events],
  155. )
  156. self.assertTrue(all(event['company_id'] == 1002 for event in reporter.events))
  157. self.assertNotIn('GWS_A', str(reporter.events))
  158. self.assertNotIn('MT_A', str(reporter.events))
  159. def test_backend_failure_emits_failed_event_without_changing_exception(self):
  160. class FailingApiClient(FakeApiClient):
  161. def call_tool(self, *args, **kwargs):
  162. raise OSError('backend secret unavailable')
  163. store = FakeSessionStore()
  164. store.sessions['GWS_A'] = {
  165. 'mcp_token': 'MT_A',
  166. 'admin_id': 88,
  167. 'company_id': 1002,
  168. }
  169. reporter = RecordingReporter()
  170. emitter = RequestDiagnosticEmitter(reporter, 'rq_a')
  171. app = PublicGatewayApp(store, FailingApiClient(), auth_client=None)
  172. with self.assertRaisesRegex(OSError, 'backend secret unavailable'):
  173. app.call_tool(
  174. 'GWS_A',
  175. 'query_order',
  176. request_id='rq_a',
  177. diagnostic_emitter=emitter,
  178. )
  179. failed = reporter.events[-1]
  180. self.assertEqual('backend_call', failed['stage'])
  181. self.assertEqual('failed', failed['status'])
  182. self.assertEqual('UNEXPECTED_EXCEPTION', failed['event_code'])
  183. self.assertNotIn('backend secret', str(failed))
  184. def test_two_employees_use_isolated_tokens(self):
  185. store = FakeSessionStore()
  186. store.sessions['GWS_A'] = {'mcp_token': 'MT_A'}
  187. store.sessions['GWS_B'] = {'mcp_token': 'MT_B'}
  188. api_client = FakeApiClient()
  189. app = PublicGatewayApp(session_store=store, api_client=api_client, auth_client=None)
  190. result_a = app.call_tool('GWS_A', 'query_order', {'keyword': 'A'}, request_id='rq_a')
  191. result_b = app.call_tool('GWS_B', 'query_order', {'keyword': 'B'}, request_id='rq_b')
  192. self.assertEqual('MT_A', result_a['data']['token_used'])
  193. self.assertEqual('MT_B', result_b['data']['token_used'])
  194. self.assertEqual('MT_A', api_client.calls[0][0])
  195. self.assertEqual('MT_B', api_client.calls[1][0])
  196. def test_public_tools_do_not_include_bind_auth_code(self):
  197. store = FakeSessionStore()
  198. store.sessions['GWS_A'] = {'mcp_token': 'MT_A'}
  199. app = PublicGatewayApp(session_store=store, api_client=FakeApiClient(), auth_client=None)
  200. tool_names = [tool['name'] for tool in app.list_tools('GWS_A')]
  201. self.assertNotIn('bind_auth_code', tool_names)
  202. self.assertIn('query_order', tool_names)
  203. self.assertIn('query_track', tool_names)
  204. self.assertIn('query_order_exact', tool_names)
  205. self.assertIn('list_order_filter_options', tool_names)
  206. def test_missing_session_returns_human_device_message(self):
  207. app = PublicGatewayApp(session_store=FakeSessionStore(), api_client=FakeApiClient(), auth_client=None)
  208. with self.assertRaises(RuntimeError) as error:
  209. app.call_tool('GWS_missing', 'query_order', {'keyword': 'A'}, request_id='rq_missing')
  210. self.assertIn('这台设备的 Workbuddy 配置已失效,请重新生成配置', str(error.exception))
  211. def test_tool_call_forwards_client_ip_to_api_client(self):
  212. store = FakeSessionStore()
  213. store.sessions['GWS_A'] = {'mcp_token': 'MT_A'}
  214. api_client = FakeApiClient()
  215. app = PublicGatewayApp(session_store=store, api_client=api_client, auth_client=None)
  216. app.call_tool('GWS_A', 'query_order', {'keyword': 'A'}, request_id='rq_a', client_ip='203.0.113.9')
  217. self.assertEqual('203.0.113.9', api_client.calls[0][5])
  218. def test_successful_tool_call_touches_gateway_session(self):
  219. store = FakeSessionStore()
  220. store.sessions['GWS_A'] = {'mcp_token': 'MT_A'}
  221. app = PublicGatewayApp(session_store=store, api_client=FakeApiClient(), auth_client=None)
  222. app.call_tool('GWS_A', 'query_order', {'keyword': 'A'}, request_id='rq_a')
  223. self.assertEqual(['GWS_A'], store.touched)
  224. class RecordingReporter:
  225. def __init__(self):
  226. self.events = []
  227. def report(self, event):
  228. self.events.append(event)
  229. return True
  230. if __name__ == '__main__':
  231. unittest.main()