import hashlib import json import unittest from unittest.mock import MagicMock, patch from services.diagnostic_event import ( RequestDiagnosticEmitter, build_diagnostic_event, ) from services.diagnostic_reporter import ( DiagnosticReporter, NullDiagnosticReporter, _http_transport, diagnostic_reporter_from_config, ) class DiagnosticReporterTest(unittest.TestCase): def test_event_builder_rejects_invalid_required_and_optional_fields(self): base = { 'request_id': 'rq_valid', 'stage': 'request_ingress', 'status': 'started', 'event_code': 'REQUEST_RECEIVED', } cases = ( ('request_id', None), ('request_id', 'bad'), ('stage', 'unknown'), ('status', 'unknown'), ('event_code', 1), ('event_code', 'A' * 65), ('event_code', 'lowercase'), ('company_id', True), ('company_id', '1'), ('company_id', 0), ('tool_code', 1), ('tool_code', 'UPPER'), ('session_credential', 1), ('response_code', 'bad-code'), ('summary_code', 'A' * 65), ('cost_ms', True), ('cost_ms', '1'), ('cost_ms', -1), ('cost_ms', 3600001), ) for field, value in cases: with self.subTest(field=field, value=value): values = dict(base) values[field] = value with self.assertRaises(ValueError): build_diagnostic_event(**values) def test_event_builder_enforces_stage_context_types(self): cases = ( ('request_ingress', 'not-object'), ('request_ingress', {'secret': 'x'}), ('request_ingress', {'http_status': True}), ('request_ingress', {'http_status': '200'}), ('request_ingress', {'http_status': 99}), ('protocol_validation', {'jsonrpc_code': True}), ('protocol_validation', {'jsonrpc_code': 'bad'}), ('response_write', {'client_disconnected': 1}), ('request_ingress', {'transport': 1}), ('request_ingress', {'transport': ''}), ('request_ingress', {'transport': 'x' * 65}), ) for stage, context in cases: with self.subTest(stage=stage, context=context): with self.assertRaises(ValueError): build_diagnostic_event( request_id='rq_valid', stage=stage, status='failed', event_code='INVALID_EVENT', context=context, ) def test_event_builder_accepts_complete_whitelisted_event(self): event = build_diagnostic_event( request_id='rq_complete', stage='response_write', status='succeeded', event_code='RESPONSE_WRITE_COMPLETED', company_id=1, admin_id=2, tool_code='query_order', response_code='MCP_0000', summary_code='OK', cost_ms=0, context={ 'http_status': 200, 'client_disconnected': False, 'transport': 'http', }, ) self.assertTrue(event['occurred_at'].endswith('Z')) self.assertEqual('OK', event['summary_code']) def test_event_builder_hashes_session_and_keeps_whitelist(self): event = build_diagnostic_event( request_id='rq_http_test_1', stage='gateway_session', status='failed', event_code='GATEWAY_SESSION_NOT_FOUND', occurred_at='2026-07-20T02:00:00.000Z', session_credential='GWS_super_secret', company_id=1002, admin_id=88, tool_code='query_order_detail', context={'transport': 'http'}, ) self.assertTrue(event['event_id'].startswith('evt_gateway_')) self.assertEqual( hashlib.sha256(b'GWS_super_secret').hexdigest()[:12], event['session_hash'], ) self.assertNotIn('session_credential', event) self.assertNotIn('GWS_super_secret', json.dumps(event)) self.assertEqual('gateway', event['source']) def test_request_emitter_caps_each_request_at_twenty_events(self): reporter = RecordingReporter() emitter = RequestDiagnosticEmitter(reporter, 'rq_http_test_1') for index in range(25): emitter.emit( stage='request_ingress', status='started', event_code='REQUEST_RECEIVED', context={'transport': 'http'}, ) self.assertEqual(20, len(reporter.events)) self.assertEqual(5, emitter.dropped_count) self.assertEqual(20, len({item['event_id'] for item in reporter.events})) def test_request_emitter_counts_reporter_rejection(self): reporter = RecordingReporter(accepted=False) emitter = RequestDiagnosticEmitter(reporter, 'rq_rejected') emitter.set_defaults(context={'transport': 'http'}) self.assertFalse(emitter.emit( stage='request_ingress', status='started', event_code='REQUEST_RECEIVED', )) self.assertEqual(1, emitter.dropped_count) def test_deferred_emitter_enriches_prestage_events_after_identity_resolution(self): reporter = RecordingReporter() emitter = RequestDiagnosticEmitter( reporter, 'rq_deferred', defer_until_identity=True, ) emitter.emit( stage='request_ingress', status='started', event_code='REQUEST_RECEIVED', context={'transport': 'http'}, ) emitter.emit( stage='protocol_validation', status='succeeded', event_code='PROTOCOL_VALIDATION_COMPLETED', context={'transport': 'http'}, ) self.assertEqual([], reporter.events) emitter.set_defaults( company_id=1002, admin_id=88, session_credential='GWS_private', ) self.assertEqual(2, len(reporter.events)) self.assertTrue(all( event['company_id'] == 1002 for event in reporter.events )) self.assertNotIn('GWS_private', str(reporter.events)) def test_deferred_emitter_flushes_unscoped_failures(self): reporter = RecordingReporter() emitter = RequestDiagnosticEmitter( reporter, 'rq_unscoped', defer_until_identity=True, ) emitter.emit( stage='gateway_session', status='failed', event_code='GATEWAY_SESSION_NOT_FOUND', session_credential='GWS_missing', context={'transport': 'http'}, ) self.assertTrue(emitter.flush()) self.assertEqual(1, len(reporter.events)) self.assertIn('session_hash', reporter.events[0]) self.assertNotIn('company_id', reporter.events[0]) def test_queue_full_only_increments_drop_count(self): reporter = self.reporter(queue_size=1) first = build_diagnostic_event( request_id='rq_http_1', stage='request_ingress', status='started', event_code='REQUEST_RECEIVED', ) second = build_diagnostic_event( request_id='rq_http_2', stage='request_ingress', status='started', event_code='REQUEST_RECEIVED', ) self.assertTrue(reporter.report(first)) self.assertFalse(reporter.report(second)) self.assertEqual(1, reporter.drop_count) def test_batch_signs_exact_json_bytes_and_accepts_success(self): calls = [] def transport(url, body, headers, timeout): calls.append((url, body, headers, timeout)) return 200, json.dumps({ 'code': 'MCP_DIAG_INGEST_0000', 'msg': 'success', 'data': {'accepted': 1, 'duplicate': 0, 'rejected': 0}, }).encode('utf-8') reporter = self.reporter(transport=transport) reporter.report(build_diagnostic_event( request_id='rq_http_1', stage='request_ingress', status='started', event_code='REQUEST_RECEIVED', )) self.assertTrue(reporter.process_once()) self.assertEqual(1, len(calls)) url, body, headers, timeout = calls[0] self.assertEqual('https://support.internal/internal/mcp-diagnostics/events', url) self.assertEqual(0.5, timeout) self.assertEqual({'events'}, set(json.loads(body.decode('utf-8')))) canonical = 'POST\n/internal/mcp-diagnostics/events\n{0}\n{1}\n{2}'.format( headers['X-MCP-Timestamp'], headers['X-MCP-Nonce'], hashlib.sha256(body).hexdigest(), ) expected = hashlib.pbkdf2_hmac( 'sha256', canonical.encode('utf-8'), b'x', 1 ) self.assertNotEqual(expected.hex(), headers['X-MCP-Signature']) import hmac self.assertEqual( hmac.new(b's' * 32, canonical.encode('utf-8'), hashlib.sha256).hexdigest(), headers['X-MCP-Signature'], ) def test_failed_send_keeps_batch_and_uses_bounded_backoff(self): attempts = [] sleeps = [] def transport(_url, _body, _headers, _timeout): attempts.append(1) if len(attempts) == 1: raise OSError('support unavailable secret') return 200, b'{"code":"MCP_DIAG_INGEST_0000","data":{}}' reporter = self.reporter( transport=transport, sleeper=sleeps.append, initial_backoff=0.25, max_backoff=1.0, ) reporter.report(build_diagnostic_event( request_id='rq_http_1', stage='request_ingress', status='started', event_code='REQUEST_RECEIVED', )) self.assertFalse(reporter.process_once()) self.assertEqual([0.25], sleeps) self.assertTrue(reporter.process_once()) self.assertEqual(2, len(attempts)) self.assertEqual(0, reporter.pending_count) def test_thread_start_failure_and_null_reporter_are_non_fatal(self): class BrokenThread: def start(self): raise RuntimeError('thread unavailable') reporter = self.reporter(thread_factory=lambda **_kwargs: BrokenThread()) self.assertFalse(reporter.start()) reporter.close() null = NullDiagnosticReporter() self.assertTrue(null.start()) self.assertTrue(null.report({'ignored': True})) null.close() def test_start_is_idempotent_and_close_joins_live_thread(self): thread = AliveThread() reporter = self.reporter(thread_factory=lambda **_kwargs: thread) self.assertTrue(reporter.start()) self.assertTrue(reporter.start()) reporter.close() self.assertEqual([0.5], thread.join_timeouts) def test_process_once_accepts_empty_queue_and_rejects_bad_responses(self): reporter = self.reporter() self.assertTrue(reporter.process_once()) responses = ( (199, b'{"code":"MCP_DIAG_INGEST_0000"}'), (300, b'{"code":"MCP_DIAG_INGEST_0000"}'), (200, b'{"code":"MCP_DIAG_INGEST_INVALID"}'), ) for status, body in responses: with self.subTest(status=status, body=body): reporter = self.reporter( transport=lambda *_args, result=(status, body): result, ) reporter.report(build_diagnostic_event( request_id='rq_rejected', stage='request_ingress', status='started', event_code='REQUEST_RECEIVED', )) self.assertFalse(reporter.process_once()) self.assertEqual(1, reporter.pending_count) def test_batch_size_caps_single_send_and_run_loop_handles_retry(self): reporter = self.reporter(batch_size=1) for request_id in ('rq_one', 'rq_two'): reporter.report(build_diagnostic_event( request_id=request_id, stage='request_ingress', status='started', event_code='REQUEST_RECEIVED', )) self.assertTrue(reporter.process_once()) self.assertEqual(0, reporter.pending_count) self.assertFalse(reporter._queue.empty()) reporter = self.reporter() reporter._stop = MagicMock() reporter._stop.is_set.side_effect = (False, False, True) reporter.process_once = MagicMock(side_effect=(False, True)) reporter._run() reporter._stop.wait.assert_called_once_with(0.1) def test_default_http_transport_posts_and_reads_response(self): response = MagicMock() response.getcode.return_value = 200 response.read.return_value = b'{}' response.__enter__.return_value = response response.__exit__.return_value = False with patch( 'services.diagnostic_reporter.urllib.request.urlopen', return_value=response, ) as urlopen: result = _http_transport( 'https://support.test/path', b'{}', {'X-Test': '1'}, 0.5, ) self.assertEqual((200, b'{}'), result) self.assertEqual(0.5, urlopen.call_args.kwargs['timeout']) def test_config_factory_uses_null_when_disabled_or_invalid(self): disabled = type('Config', (), {'diagnosis_enabled': False})() invalid = type('Config', (), { 'diagnosis_enabled': True, 'diagnosis_url': '', 'diagnosis_key_id': '', 'diagnosis_secret': '', })() self.assertIsInstance( diagnostic_reporter_from_config(disabled), NullDiagnosticReporter, ) self.assertIsInstance( diagnostic_reporter_from_config(invalid), NullDiagnosticReporter, ) def test_config_factory_starts_enabled_reporter(self): config = type('Config', (), { 'diagnosis_enabled': True, 'diagnosis_url': 'https://support.test/internal/mcp-diagnostics/events', 'diagnosis_key_id': 'gateway-current', 'diagnosis_secret': 's' * 32, 'diagnosis_queue_size': 10, 'diagnosis_batch_size': 20, 'diagnosis_timeout_seconds': 0.5, 'diagnosis_initial_backoff_seconds': 0.25, 'diagnosis_max_backoff_seconds': 2.0, })() reporter = diagnostic_reporter_from_config( config, thread_factory=lambda **_kwargs: StartedThread(), ) self.assertIsInstance(reporter, DiagnosticReporter) self.assertTrue(reporter._thread.started) reporter.close() def test_config_factory_requires_explicit_opt_in_for_http(self): values = { 'diagnosis_enabled': True, 'diagnosis_url': 'http://support.test/internal/mcp-diagnostics/events', 'diagnosis_key_id': 'gateway-current', 'diagnosis_secret': 's' * 32, } default_config = type('Config', (), values)() allowed_config = type( 'Config', (), dict(values, diagnosis_allow_insecure_http=True), )() self.assertIsInstance( diagnostic_reporter_from_config(default_config), NullDiagnosticReporter, ) reporter = diagnostic_reporter_from_config( allowed_config, thread_factory=lambda **_kwargs: StartedThread(), ) self.assertIsInstance(reporter, DiagnosticReporter) reporter.close() def test_config_factory_falls_back_when_thread_cannot_start(self): config = type('Config', (), { 'diagnosis_enabled': True, 'diagnosis_url': 'https://support.test/events', 'diagnosis_key_id': 'gateway-current', 'diagnosis_secret': 's' * 32, })() reporter = diagnostic_reporter_from_config( config, thread_factory=lambda **_kwargs: BrokenThread(), ) self.assertIsInstance(reporter, NullDiagnosticReporter) def reporter(self, **overrides): options = { 'url': 'https://support.internal/internal/mcp-diagnostics/events', 'key_id': 'gateway-current', 'secret': 's' * 32, 'queue_size': 10, 'batch_size': 100, 'timeout_seconds': 0.5, 'transport': lambda *_args: (200, b'{"code":"MCP_DIAG_INGEST_0000","data":{}}'), 'clock': lambda: 1784512800, 'nonce_factory': lambda: 'nonce-gateway-1234', 'sleeper': lambda _seconds: None, } options.update(overrides) return DiagnosticReporter(**options) class RecordingReporter: def __init__(self, accepted=True): self.events = [] self.accepted = accepted def report(self, event): self.events.append(event) return self.accepted class StartedThread: def __init__(self): self.started = False def start(self): self.started = True def is_alive(self): return False class AliveThread(StartedThread): def __init__(self): super().__init__() self.join_timeouts = [] def is_alive(self): return True def join(self, timeout=None): self.join_timeouts.append(timeout) class BrokenThread: def start(self): raise RuntimeError('thread unavailable') if __name__ == '__main__': unittest.main()