test_diagnostic_reporter.py 18 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510
  1. import hashlib
  2. import json
  3. import unittest
  4. from unittest.mock import MagicMock, patch
  5. from services.diagnostic_event import (
  6. RequestDiagnosticEmitter,
  7. build_diagnostic_event,
  8. )
  9. from services.diagnostic_reporter import (
  10. DiagnosticReporter,
  11. NullDiagnosticReporter,
  12. _http_transport,
  13. diagnostic_reporter_from_config,
  14. )
  15. class DiagnosticReporterTest(unittest.TestCase):
  16. def test_event_builder_rejects_invalid_required_and_optional_fields(self):
  17. base = {
  18. 'request_id': 'rq_valid',
  19. 'stage': 'request_ingress',
  20. 'status': 'started',
  21. 'event_code': 'REQUEST_RECEIVED',
  22. }
  23. cases = (
  24. ('request_id', None),
  25. ('request_id', 'bad'),
  26. ('stage', 'unknown'),
  27. ('status', 'unknown'),
  28. ('event_code', 1),
  29. ('event_code', 'A' * 65),
  30. ('event_code', 'lowercase'),
  31. ('company_id', True),
  32. ('company_id', '1'),
  33. ('company_id', 0),
  34. ('tool_code', 1),
  35. ('tool_code', 'UPPER'),
  36. ('session_credential', 1),
  37. ('response_code', 'bad-code'),
  38. ('summary_code', 'A' * 65),
  39. ('cost_ms', True),
  40. ('cost_ms', '1'),
  41. ('cost_ms', -1),
  42. ('cost_ms', 3600001),
  43. )
  44. for field, value in cases:
  45. with self.subTest(field=field, value=value):
  46. values = dict(base)
  47. values[field] = value
  48. with self.assertRaises(ValueError):
  49. build_diagnostic_event(**values)
  50. def test_event_builder_enforces_stage_context_types(self):
  51. cases = (
  52. ('request_ingress', 'not-object'),
  53. ('request_ingress', {'secret': 'x'}),
  54. ('request_ingress', {'http_status': True}),
  55. ('request_ingress', {'http_status': '200'}),
  56. ('request_ingress', {'http_status': 99}),
  57. ('protocol_validation', {'jsonrpc_code': True}),
  58. ('protocol_validation', {'jsonrpc_code': 'bad'}),
  59. ('response_write', {'client_disconnected': 1}),
  60. ('request_ingress', {'transport': 1}),
  61. ('request_ingress', {'transport': ''}),
  62. ('request_ingress', {'transport': 'x' * 65}),
  63. )
  64. for stage, context in cases:
  65. with self.subTest(stage=stage, context=context):
  66. with self.assertRaises(ValueError):
  67. build_diagnostic_event(
  68. request_id='rq_valid',
  69. stage=stage,
  70. status='failed',
  71. event_code='INVALID_EVENT',
  72. context=context,
  73. )
  74. def test_event_builder_accepts_complete_whitelisted_event(self):
  75. event = build_diagnostic_event(
  76. request_id='rq_complete',
  77. stage='response_write',
  78. status='succeeded',
  79. event_code='RESPONSE_WRITE_COMPLETED',
  80. company_id=1,
  81. admin_id=2,
  82. tool_code='query_order',
  83. response_code='MCP_0000',
  84. summary_code='OK',
  85. cost_ms=0,
  86. context={
  87. 'http_status': 200,
  88. 'client_disconnected': False,
  89. 'transport': 'http',
  90. },
  91. )
  92. self.assertTrue(event['occurred_at'].endswith('Z'))
  93. self.assertEqual('OK', event['summary_code'])
  94. def test_event_builder_hashes_session_and_keeps_whitelist(self):
  95. event = build_diagnostic_event(
  96. request_id='rq_http_test_1',
  97. stage='gateway_session',
  98. status='failed',
  99. event_code='GATEWAY_SESSION_NOT_FOUND',
  100. occurred_at='2026-07-20T02:00:00.000Z',
  101. session_credential='GWS_super_secret',
  102. company_id=1002,
  103. admin_id=88,
  104. tool_code='query_order_detail',
  105. context={'transport': 'http'},
  106. )
  107. self.assertTrue(event['event_id'].startswith('evt_gateway_'))
  108. self.assertEqual(
  109. hashlib.sha256(b'GWS_super_secret').hexdigest()[:12],
  110. event['session_hash'],
  111. )
  112. self.assertNotIn('session_credential', event)
  113. self.assertNotIn('GWS_super_secret', json.dumps(event))
  114. self.assertEqual('gateway', event['source'])
  115. def test_request_emitter_caps_each_request_at_twenty_events(self):
  116. reporter = RecordingReporter()
  117. emitter = RequestDiagnosticEmitter(reporter, 'rq_http_test_1')
  118. for index in range(25):
  119. emitter.emit(
  120. stage='request_ingress',
  121. status='started',
  122. event_code='REQUEST_RECEIVED',
  123. context={'transport': 'http'},
  124. )
  125. self.assertEqual(20, len(reporter.events))
  126. self.assertEqual(5, emitter.dropped_count)
  127. self.assertEqual(20, len({item['event_id'] for item in reporter.events}))
  128. def test_request_emitter_counts_reporter_rejection(self):
  129. reporter = RecordingReporter(accepted=False)
  130. emitter = RequestDiagnosticEmitter(reporter, 'rq_rejected')
  131. emitter.set_defaults(context={'transport': 'http'})
  132. self.assertFalse(emitter.emit(
  133. stage='request_ingress',
  134. status='started',
  135. event_code='REQUEST_RECEIVED',
  136. ))
  137. self.assertEqual(1, emitter.dropped_count)
  138. def test_deferred_emitter_enriches_prestage_events_after_identity_resolution(self):
  139. reporter = RecordingReporter()
  140. emitter = RequestDiagnosticEmitter(
  141. reporter,
  142. 'rq_deferred',
  143. defer_until_identity=True,
  144. )
  145. emitter.emit(
  146. stage='request_ingress',
  147. status='started',
  148. event_code='REQUEST_RECEIVED',
  149. context={'transport': 'http'},
  150. )
  151. emitter.emit(
  152. stage='protocol_validation',
  153. status='succeeded',
  154. event_code='PROTOCOL_VALIDATION_COMPLETED',
  155. context={'transport': 'http'},
  156. )
  157. self.assertEqual([], reporter.events)
  158. emitter.set_defaults(
  159. company_id=1002,
  160. admin_id=88,
  161. session_credential='GWS_private',
  162. )
  163. self.assertEqual(2, len(reporter.events))
  164. self.assertTrue(all(
  165. event['company_id'] == 1002 for event in reporter.events
  166. ))
  167. self.assertNotIn('GWS_private', str(reporter.events))
  168. def test_deferred_emitter_flushes_unscoped_failures(self):
  169. reporter = RecordingReporter()
  170. emitter = RequestDiagnosticEmitter(
  171. reporter,
  172. 'rq_unscoped',
  173. defer_until_identity=True,
  174. )
  175. emitter.emit(
  176. stage='gateway_session',
  177. status='failed',
  178. event_code='GATEWAY_SESSION_NOT_FOUND',
  179. session_credential='GWS_missing',
  180. context={'transport': 'http'},
  181. )
  182. self.assertTrue(emitter.flush())
  183. self.assertEqual(1, len(reporter.events))
  184. self.assertIn('session_hash', reporter.events[0])
  185. self.assertNotIn('company_id', reporter.events[0])
  186. def test_queue_full_only_increments_drop_count(self):
  187. reporter = self.reporter(queue_size=1)
  188. first = build_diagnostic_event(
  189. request_id='rq_http_1', stage='request_ingress',
  190. status='started', event_code='REQUEST_RECEIVED',
  191. )
  192. second = build_diagnostic_event(
  193. request_id='rq_http_2', stage='request_ingress',
  194. status='started', event_code='REQUEST_RECEIVED',
  195. )
  196. self.assertTrue(reporter.report(first))
  197. self.assertFalse(reporter.report(second))
  198. self.assertEqual(1, reporter.drop_count)
  199. def test_batch_signs_exact_json_bytes_and_accepts_success(self):
  200. calls = []
  201. def transport(url, body, headers, timeout):
  202. calls.append((url, body, headers, timeout))
  203. return 200, json.dumps({
  204. 'code': 'MCP_DIAG_INGEST_0000',
  205. 'msg': 'success',
  206. 'data': {'accepted': 1, 'duplicate': 0, 'rejected': 0},
  207. }).encode('utf-8')
  208. reporter = self.reporter(transport=transport)
  209. reporter.report(build_diagnostic_event(
  210. request_id='rq_http_1', stage='request_ingress',
  211. status='started', event_code='REQUEST_RECEIVED',
  212. ))
  213. self.assertTrue(reporter.process_once())
  214. self.assertEqual(1, len(calls))
  215. url, body, headers, timeout = calls[0]
  216. self.assertEqual('https://support.internal/internal/mcp-diagnostics/events', url)
  217. self.assertEqual(0.5, timeout)
  218. self.assertEqual({'events'}, set(json.loads(body.decode('utf-8'))))
  219. canonical = 'POST\n/internal/mcp-diagnostics/events\n{0}\n{1}\n{2}'.format(
  220. headers['X-MCP-Timestamp'],
  221. headers['X-MCP-Nonce'],
  222. hashlib.sha256(body).hexdigest(),
  223. )
  224. expected = hashlib.pbkdf2_hmac(
  225. 'sha256', canonical.encode('utf-8'), b'x', 1
  226. )
  227. self.assertNotEqual(expected.hex(), headers['X-MCP-Signature'])
  228. import hmac
  229. self.assertEqual(
  230. hmac.new(b's' * 32, canonical.encode('utf-8'), hashlib.sha256).hexdigest(),
  231. headers['X-MCP-Signature'],
  232. )
  233. def test_failed_send_keeps_batch_and_uses_bounded_backoff(self):
  234. attempts = []
  235. sleeps = []
  236. def transport(_url, _body, _headers, _timeout):
  237. attempts.append(1)
  238. if len(attempts) == 1:
  239. raise OSError('support unavailable secret')
  240. return 200, b'{"code":"MCP_DIAG_INGEST_0000","data":{}}'
  241. reporter = self.reporter(
  242. transport=transport,
  243. sleeper=sleeps.append,
  244. initial_backoff=0.25,
  245. max_backoff=1.0,
  246. )
  247. reporter.report(build_diagnostic_event(
  248. request_id='rq_http_1', stage='request_ingress',
  249. status='started', event_code='REQUEST_RECEIVED',
  250. ))
  251. self.assertFalse(reporter.process_once())
  252. self.assertEqual([0.25], sleeps)
  253. self.assertTrue(reporter.process_once())
  254. self.assertEqual(2, len(attempts))
  255. self.assertEqual(0, reporter.pending_count)
  256. def test_thread_start_failure_and_null_reporter_are_non_fatal(self):
  257. class BrokenThread:
  258. def start(self):
  259. raise RuntimeError('thread unavailable')
  260. reporter = self.reporter(thread_factory=lambda **_kwargs: BrokenThread())
  261. self.assertFalse(reporter.start())
  262. reporter.close()
  263. null = NullDiagnosticReporter()
  264. self.assertTrue(null.start())
  265. self.assertTrue(null.report({'ignored': True}))
  266. null.close()
  267. def test_start_is_idempotent_and_close_joins_live_thread(self):
  268. thread = AliveThread()
  269. reporter = self.reporter(thread_factory=lambda **_kwargs: thread)
  270. self.assertTrue(reporter.start())
  271. self.assertTrue(reporter.start())
  272. reporter.close()
  273. self.assertEqual([0.5], thread.join_timeouts)
  274. def test_process_once_accepts_empty_queue_and_rejects_bad_responses(self):
  275. reporter = self.reporter()
  276. self.assertTrue(reporter.process_once())
  277. responses = (
  278. (199, b'{"code":"MCP_DIAG_INGEST_0000"}'),
  279. (300, b'{"code":"MCP_DIAG_INGEST_0000"}'),
  280. (200, b'{"code":"MCP_DIAG_INGEST_INVALID"}'),
  281. )
  282. for status, body in responses:
  283. with self.subTest(status=status, body=body):
  284. reporter = self.reporter(
  285. transport=lambda *_args, result=(status, body): result,
  286. )
  287. reporter.report(build_diagnostic_event(
  288. request_id='rq_rejected',
  289. stage='request_ingress',
  290. status='started',
  291. event_code='REQUEST_RECEIVED',
  292. ))
  293. self.assertFalse(reporter.process_once())
  294. self.assertEqual(1, reporter.pending_count)
  295. def test_batch_size_caps_single_send_and_run_loop_handles_retry(self):
  296. reporter = self.reporter(batch_size=1)
  297. for request_id in ('rq_one', 'rq_two'):
  298. reporter.report(build_diagnostic_event(
  299. request_id=request_id,
  300. stage='request_ingress',
  301. status='started',
  302. event_code='REQUEST_RECEIVED',
  303. ))
  304. self.assertTrue(reporter.process_once())
  305. self.assertEqual(0, reporter.pending_count)
  306. self.assertFalse(reporter._queue.empty())
  307. reporter = self.reporter()
  308. reporter._stop = MagicMock()
  309. reporter._stop.is_set.side_effect = (False, False, True)
  310. reporter.process_once = MagicMock(side_effect=(False, True))
  311. reporter._run()
  312. reporter._stop.wait.assert_called_once_with(0.1)
  313. def test_default_http_transport_posts_and_reads_response(self):
  314. response = MagicMock()
  315. response.getcode.return_value = 200
  316. response.read.return_value = b'{}'
  317. response.__enter__.return_value = response
  318. response.__exit__.return_value = False
  319. with patch(
  320. 'services.diagnostic_reporter.urllib.request.urlopen',
  321. return_value=response,
  322. ) as urlopen:
  323. result = _http_transport(
  324. 'https://support.test/path',
  325. b'{}',
  326. {'X-Test': '1'},
  327. 0.5,
  328. )
  329. self.assertEqual((200, b'{}'), result)
  330. self.assertEqual(0.5, urlopen.call_args.kwargs['timeout'])
  331. def test_config_factory_uses_null_when_disabled_or_invalid(self):
  332. disabled = type('Config', (), {'diagnosis_enabled': False})()
  333. invalid = type('Config', (), {
  334. 'diagnosis_enabled': True,
  335. 'diagnosis_url': '',
  336. 'diagnosis_key_id': '',
  337. 'diagnosis_secret': '',
  338. })()
  339. self.assertIsInstance(
  340. diagnostic_reporter_from_config(disabled),
  341. NullDiagnosticReporter,
  342. )
  343. self.assertIsInstance(
  344. diagnostic_reporter_from_config(invalid),
  345. NullDiagnosticReporter,
  346. )
  347. def test_config_factory_starts_enabled_reporter(self):
  348. config = type('Config', (), {
  349. 'diagnosis_enabled': True,
  350. 'diagnosis_url': 'https://support.test/internal/mcp-diagnostics/events',
  351. 'diagnosis_key_id': 'gateway-current',
  352. 'diagnosis_secret': 's' * 32,
  353. 'diagnosis_queue_size': 10,
  354. 'diagnosis_batch_size': 20,
  355. 'diagnosis_timeout_seconds': 0.5,
  356. 'diagnosis_initial_backoff_seconds': 0.25,
  357. 'diagnosis_max_backoff_seconds': 2.0,
  358. })()
  359. reporter = diagnostic_reporter_from_config(
  360. config,
  361. thread_factory=lambda **_kwargs: StartedThread(),
  362. )
  363. self.assertIsInstance(reporter, DiagnosticReporter)
  364. self.assertTrue(reporter._thread.started)
  365. reporter.close()
  366. def test_config_factory_requires_explicit_opt_in_for_http(self):
  367. values = {
  368. 'diagnosis_enabled': True,
  369. 'diagnosis_url': 'http://support.test/internal/mcp-diagnostics/events',
  370. 'diagnosis_key_id': 'gateway-current',
  371. 'diagnosis_secret': 's' * 32,
  372. }
  373. default_config = type('Config', (), values)()
  374. allowed_config = type(
  375. 'Config',
  376. (),
  377. dict(values, diagnosis_allow_insecure_http=True),
  378. )()
  379. self.assertIsInstance(
  380. diagnostic_reporter_from_config(default_config),
  381. NullDiagnosticReporter,
  382. )
  383. reporter = diagnostic_reporter_from_config(
  384. allowed_config,
  385. thread_factory=lambda **_kwargs: StartedThread(),
  386. )
  387. self.assertIsInstance(reporter, DiagnosticReporter)
  388. reporter.close()
  389. def test_config_factory_falls_back_when_thread_cannot_start(self):
  390. config = type('Config', (), {
  391. 'diagnosis_enabled': True,
  392. 'diagnosis_url': 'https://support.test/events',
  393. 'diagnosis_key_id': 'gateway-current',
  394. 'diagnosis_secret': 's' * 32,
  395. })()
  396. reporter = diagnostic_reporter_from_config(
  397. config,
  398. thread_factory=lambda **_kwargs: BrokenThread(),
  399. )
  400. self.assertIsInstance(reporter, NullDiagnosticReporter)
  401. def reporter(self, **overrides):
  402. options = {
  403. 'url': 'https://support.internal/internal/mcp-diagnostics/events',
  404. 'key_id': 'gateway-current',
  405. 'secret': 's' * 32,
  406. 'queue_size': 10,
  407. 'batch_size': 100,
  408. 'timeout_seconds': 0.5,
  409. 'transport': lambda *_args: (200, b'{"code":"MCP_DIAG_INGEST_0000","data":{}}'),
  410. 'clock': lambda: 1784512800,
  411. 'nonce_factory': lambda: 'nonce-gateway-1234',
  412. 'sleeper': lambda _seconds: None,
  413. }
  414. options.update(overrides)
  415. return DiagnosticReporter(**options)
  416. class RecordingReporter:
  417. def __init__(self, accepted=True):
  418. self.events = []
  419. self.accepted = accepted
  420. def report(self, event):
  421. self.events.append(event)
  422. return self.accepted
  423. class StartedThread:
  424. def __init__(self):
  425. self.started = False
  426. def start(self):
  427. self.started = True
  428. def is_alive(self):
  429. return False
  430. class AliveThread(StartedThread):
  431. def __init__(self):
  432. super().__init__()
  433. self.join_timeouts = []
  434. def is_alive(self):
  435. return True
  436. def join(self, timeout=None):
  437. self.join_timeouts.append(timeout)
  438. class BrokenThread:
  439. def start(self):
  440. raise RuntimeError('thread unavailable')
  441. if __name__ == '__main__':
  442. unittest.main()