import dataclasses import hashlib import json import os import sys import unittest APP_DIR = os.path.abspath(os.path.join(os.path.dirname(__file__), '..', 'app')) if APP_DIR not in sys.path: sys.path.insert(0, APP_DIR) from worker_contracts import ( ALLOWED_PHASE_TRANSITIONS, CANONICAL_WORKER_PHASES, DIAGNOSTIC_SCHEMA, MAX_DIAGNOSTIC_AGGREGATE_BYTES, MAX_DIAGNOSTIC_BODY_BYTES, MAX_DIAGNOSTIC_ENVELOPE_BYTES, MAX_DIAGNOSTIC_LOG_BYTES, MAX_DIAGNOSTICS_PER_ASSIGNMENT, WORKER_EVENT_SCHEMA, WORKER_EVENT_TYPE, AssignmentOutcome, DiagnosticCategory, DiagnosticEnvelope, DiagnosticExceptionContext, DiagnosticHTTPContext, DiagnosticKind, DiagnosticMaterial, DiagnosticProcessContext, MaterialEncoding, ScanOutcome, WorkerContractError, WorkerEvent, WorkerPhase, build_diagnostic_envelope, build_legacy_error_frame_diagnostics, decode_diagnostic_envelope, decode_diagnostic_envelopes_ndjson, decode_worker_event, decode_worker_events_ndjson, diagnostic_material_bytes, encode_diagnostic_envelope, encode_diagnostic_envelopes_ndjson, encode_worker_event, encode_worker_events_ndjson, make_body_material, make_diagnostic_material, make_log_material, validate_diagnostic_envelopes, validate_phase_transition, validate_worker_event_sequence, ) UTC = '2026-09-23T12:00:00Z' class WorkerContractTests(unittest.TestCase): def event(self, sequence=1, phase=WorkerPhase.IDLE, **changes): values = { 'schema': WORKER_EVENT_SCHEMA, 'sequence': sequence, 'timestamp': UTC, 'instance_id': 'worker-instance-1', 'slot_id': 0, 'reservation_id': None, 'source': None, 'type': WORKER_EVENT_TYPE, 'phase': phase, 'phase_started_at': UTC, 'scan_deadline_at': None, 'assignment_deadline_at': None, 'progress': {}, } values.update(changes) return WorkerEvent(**values) def envelope(self, **changes): values = { 'occurrence_id': 'runner-http-1', 'reservation_id': 123, 'scan_event_id': 'b' * 32, 'slot_id': 0, 'source': 'dockerhub', 'phase': WorkerPhase.RESOLVING, 'kind': DiagnosticKind.PROVIDER_HTTP, 'category': DiagnosticCategory.AUTHORIZATION, 'code': 'docker.manifest_http_403', 'summary': 'manifest request was denied', 'retryable': False, 'attempt': 1, 'assignment_outcome': AssignmentOutcome.ACCEPTED, 'scan_outcome': ScanOutcome.ERROR, 'occurred_at': UTC, 'captured_at': '2026-09-23T12:00:00.100000Z', 'received_at': None, 'http': DiagnosticHTTPContext( operation='manifest.get', status_code=403, content_type='application/json; charset=utf-8', request_id='request-1', body=make_body_material(b'{"error":"denied"}'), ), 'process': DiagnosticProcessContext( name='trufflehog', exit_code=1, signal=None, timed_out=False, stdout=make_log_material(b'stdout line\n', maximum=1024), stderr=make_log_material(b'stderr line\n', maximum=1024), ), 'exception': DiagnosticExceptionContext( type='ProviderError', message='request failed', fingerprint='provider-error-v1', ), } values.update(changes) return build_diagnostic_envelope(**values) def test_canonical_phase_vocabulary_and_transition_edges_are_exact(self): self.assertEqual( CANONICAL_WORKER_PHASES, ( 'idle', 'claiming', 'assigned', 'waiting_permit', 'preparing', 'resolving', 'downloading', 'cloning', 'scanning', 'filtering', 'cleaning', 'bundling', 'uploading', 'awaiting_receipt', 'backoff', 'draining', 'stopped', ), ) self.assertEqual(set(ALLOWED_PHASE_TRANSITIONS), set(WorkerPhase)) normal = ( WorkerPhase.IDLE, WorkerPhase.CLAIMING, WorkerPhase.ASSIGNED, WorkerPhase.PREPARING, WorkerPhase.WAITING_PERMIT, WorkerPhase.RESOLVING, WorkerPhase.DOWNLOADING, WorkerPhase.SCANNING, WorkerPhase.FILTERING, WorkerPhase.CLEANING, WorkerPhase.BUNDLING, WorkerPhase.UPLOADING, WorkerPhase.AWAITING_RECEIPT, WorkerPhase.IDLE, WorkerPhase.DRAINING, WorkerPhase.STOPPED, ) for previous, current in zip(normal, normal[1:]): self.assertEqual(validate_phase_transition(previous, current), current) self.assertEqual( validate_phase_transition(WorkerPhase.SCANNING, WorkerPhase.SCANNING), WorkerPhase.SCANNING, ) for provider_phase in ( WorkerPhase.RESOLVING, WorkerPhase.DOWNLOADING, WorkerPhase.CLONING, ): self.assertEqual( validate_phase_transition(WorkerPhase.SCANNING, provider_phase), provider_phase, ) for interrupted_provider_phase in ( WorkerPhase.RESOLVING, WorkerPhase.DOWNLOADING, ): self.assertEqual( validate_phase_transition( interrupted_provider_phase, WorkerPhase.FILTERING, ), WorkerPhase.FILTERING, ) with self.assertRaisesRegex(WorkerContractError, 'invalid'): validate_phase_transition(WorkerPhase.SCANNING, WorkerPhase.CLAIMING) with self.assertRaises(WorkerContractError): validate_phase_transition(WorkerPhase.STOPPED, WorkerPhase.IDLE) def test_worker_event_json_and_ndjson_are_canonical_round_trips(self): events = ( self.event(7), self.event(9, WorkerPhase.CLAIMING, progress={'attempt': 1}), self.event( 12, WorkerPhase.ASSIGNED, reservation_id=123, source='dockerhub', scan_deadline_at='2026-09-23T12:10:00Z', assignment_deadline_at='2026-09-23T14:00:00Z', ), ) payload = encode_worker_event(events[-1]) self.assertEqual( payload, json.dumps( json.loads(payload), ensure_ascii=True, sort_keys=True, separators=(',', ':'), allow_nan=False, ).encode('ascii'), ) self.assertEqual(decode_worker_event(payload), events[-1]) ndjson = encode_worker_events_ndjson(events) self.assertTrue(ndjson.endswith(b'\n')) self.assertEqual(decode_worker_events_ndjson(ndjson), events) self.assertEqual(decode_worker_events_ndjson(b''), ()) def test_worker_event_sequence_is_global_and_phase_validation_is_per_slot(self): events = ( self.event(1, slot_id=0), self.event(2, slot_id=1), self.event(3, WorkerPhase.CLAIMING, slot_id=0), self.event(4, WorkerPhase.CLAIMING, slot_id=1), ) self.assertEqual(validate_worker_event_sequence(events), events) with self.assertRaisesRegex(WorkerContractError, 'invalid'): validate_worker_event_sequence((events[0], dataclasses.replace(events[1], sequence=1))) with self.assertRaises(WorkerContractError): validate_worker_event_sequence((self.event(1, WorkerPhase.SCANNING), self.event(2, WorkerPhase.CLAIMING))) with self.assertRaises(WorkerContractError): decode_worker_events_ndjson(encode_worker_events_ndjson(events), previous_sequence=4) def test_worker_event_rejects_unknown_noncanonical_duplicate_and_bad_utc_fields(self): payload = encode_worker_event(self.event()) with self.assertRaisesRegex(WorkerContractError, 'invalid'): decode_worker_event(payload + b' ') value = json.loads(payload) value['unknown'] = None with self.assertRaises(WorkerContractError): decode_worker_event(json.dumps(value, sort_keys=True, separators=(',', ':')).encode('ascii')) duplicate = payload[:-1] + b',"schema":1}' with self.assertRaises(WorkerContractError): decode_worker_event(duplicate) for timestamp in ( '2026-09-23T12:00:00+00:00', '2026-09-23 12:00:00Z', '2026-02-30T12:00:00Z', ): with self.subTest(timestamp=timestamp): with self.assertRaises(WorkerContractError): encode_worker_event(self.event(timestamp=timestamp)) with self.assertRaises(WorkerContractError): encode_worker_event(self.event(schema=True)) with self.assertRaises(WorkerContractError): decode_worker_events_ndjson(payload) def test_material_preserves_text_binary_body_and_log_head_tail(self): text = make_body_material('snowman: \u2603'.encode('utf-8')) self.assertEqual(text.encoding, MaterialEncoding.TEXT) self.assertEqual(diagnostic_material_bytes(text), 'snowman: \u2603'.encode('utf-8')) binary = make_body_material(b'\xff\x00\xfe') self.assertEqual(binary.encoding, MaterialEncoding.BASE64) self.assertEqual(diagnostic_material_bytes(binary), b'\xff\x00\xfe') raw_body = b'x' * (MAX_DIAGNOSTIC_BODY_BYTES + 17) body = make_body_material(raw_body) self.assertTrue(body.truncated) self.assertEqual(body.original_size, len(raw_body)) self.assertEqual(body.stored_size, MAX_DIAGNOSTIC_BODY_BYTES) self.assertEqual(body.sha256, hashlib.sha256(raw_body).hexdigest()) self.assertIsNone(body.tail) raw_log = bytes(range(256)) * 200 log = make_log_material(raw_log) stored = diagnostic_material_bytes(log) half = (MAX_DIAGNOSTIC_LOG_BYTES + 1) // 2 self.assertEqual(stored, raw_log[:half] + raw_log[-(MAX_DIAGNOSTIC_LOG_BYTES - half):]) self.assertTrue(log.truncated) self.assertEqual(log.stored_size, MAX_DIAGNOSTIC_LOG_BYTES) def test_diagnostic_json_and_ndjson_round_trip_all_nested_contexts(self): envelope = self.envelope() payload = encode_diagnostic_envelope(envelope) self.assertLessEqual(len(payload), MAX_DIAGNOSTIC_ENVELOPE_BYTES) self.assertEqual(decode_diagnostic_envelope(payload), envelope) ndjson = encode_diagnostic_envelopes_ndjson((envelope,)) self.assertEqual(decode_diagnostic_envelopes_ndjson(ndjson), (envelope,)) self.assertEqual(decode_diagnostic_envelopes_ndjson(b''), ()) self.assertEqual(envelope.schema, DIAGNOSTIC_SCHEMA) self.assertEqual(envelope.assignment_outcome, AssignmentOutcome.ACCEPTED) self.assertEqual(envelope.scan_outcome, ScanOutcome.ERROR) def test_diagnostic_uid_uses_content_and_explicit_occurrence_but_not_receive_time(self): first = self.envelope(received_at=None) replay = self.envelope(received_at='2026-09-23T12:00:01Z') other_occurrence = self.envelope(occurrence_id='runner-http-2') other_content = self.envelope(summary='different summary') self.assertEqual(first.diagnostic_uid, replay.diagnostic_uid) self.assertNotEqual(first.diagnostic_uid, other_occurrence.diagnostic_uid) self.assertNotEqual(first.diagnostic_uid, other_content.diagnostic_uid) self.assertEqual(len(first.diagnostic_uid), 64) def test_diagnostic_rejects_unknown_fields_size_hash_uid_and_context_mismatches(self): envelope = self.envelope() for scan_event_id in (456, 'ABCDEF' * 6, 'f' * 31, 'g' * 32): with self.subTest(scan_event_id=scan_event_id), self.assertRaises(WorkerContractError): self.envelope(scan_event_id=scan_event_id) value = json.loads(encode_diagnostic_envelope(envelope)) value['unknown'] = True with self.assertRaises(WorkerContractError): decode_diagnostic_envelope(json.dumps(value, sort_keys=True, separators=(',', ':')).encode('ascii')) value = json.loads(encode_diagnostic_envelope(envelope)) value['http']['body']['stored_size'] += 1 with self.assertRaises(WorkerContractError): decode_diagnostic_envelope(json.dumps(value, sort_keys=True, separators=(',', ':')).encode('ascii')) value = json.loads(encode_diagnostic_envelope(envelope)) value['http']['body']['sha256'] = '0' * 64 with self.assertRaises(WorkerContractError): decode_diagnostic_envelope(json.dumps(value, sort_keys=True, separators=(',', ':')).encode('ascii')) value = json.loads(encode_diagnostic_envelope(envelope)) value['diagnostic_uid'] = 'f' * 64 with self.assertRaises(WorkerContractError): decode_diagnostic_envelope(json.dumps(value, sort_keys=True, separators=(',', ':')).encode('ascii')) with self.assertRaises(WorkerContractError): self.envelope(kind=DiagnosticKind.EXCEPTION, exception=None) def test_diagnostic_rejects_whitespace_nul_and_index_invalid_identity(self): for values in ( {'source': ' '}, {'source': 'github\x00other'}, {'source': 'x' * 65}, {'code': ' '}, {'code': 'scan bad code'}, {'code': 'x' * 257}, {'summary': '\t\r\n'}, {'summary': 'failure\x00detail'}, {'occurrence_id': 'x' * 513}, ): with self.subTest(values=values), self.assertRaises(WorkerContractError): self.envelope(**values) def test_diagnostic_body_and_combined_log_bounds_are_enforced(self): oversized_body = make_diagnostic_material( b'x' * (MAX_DIAGNOSTIC_BODY_BYTES + 1), maximum=MAX_DIAGNOSTIC_BODY_BYTES + 1, ) with self.assertRaises(WorkerContractError): self.envelope(http=dataclasses.replace(self.envelope().http, body=oversized_body)) stdout = make_log_material(b'a' * 20000, maximum=20000) stderr = make_log_material(b'b' * 20000, maximum=20000) process = DiagnosticProcessContext( name='scanner', exit_code=1, signal=None, timed_out=False, stdout=stdout, stderr=stderr, ) with self.assertRaises(WorkerContractError): self.envelope( kind=DiagnosticKind.SCANNER_PROCESS, category=DiagnosticCategory.SCANNER, http=None, process=process, ) def test_diagnostic_envelope_count_and_aggregate_bounds_are_enforced(self): oversized = self.envelope(summary='x' * MAX_DIAGNOSTIC_ENVELOPE_BYTES) with self.assertRaises(WorkerContractError): encode_diagnostic_envelope(oversized) envelope = self.envelope() with self.assertRaises(WorkerContractError): validate_diagnostic_envelopes( (envelope,) * (MAX_DIAGNOSTICS_PER_ASSIGNMENT + 1) ) large = tuple( self.envelope(occurrence_id=f'large-{index}', summary='x' * 55000) for index in range(5) ) self.assertTrue(all( len(encode_diagnostic_envelope(item)) < MAX_DIAGNOSTIC_ENVELOPE_BYTES for item in large )) with self.assertRaises(WorkerContractError): encode_diagnostic_envelopes_ndjson(large) self.assertLess(MAX_DIAGNOSTIC_AGGREGATE_BYTES, sum( len(encode_diagnostic_envelope(item)) + 1 for item in large )) def test_legacy_error_projection_is_bounded_deterministic_and_aggregated(self): errors = [f'legacy error {index}' for index in range(40)] first = build_legacy_error_frame_diagnostics( reservation_id=7, scan_event_id='a' * 32, slot_id=0, source='github', timestamp='2026-09-24T00:00:00+00:00', errors=errors, ) replay = build_legacy_error_frame_diagnostics( reservation_id=7, scan_event_id='a' * 32, slot_id=0, source='github', timestamp='2026-09-24T00:00:00+00:00', errors=errors, ) self.assertEqual(len(first), 32) self.assertEqual(first, replay) self.assertEqual(first[-1].code, 'legacy.error_frame_aggregate') self.assertEqual(first[-1].process.stderr.original_size, sum( len(error.encode('utf-8')) + 1 for error in errors[31:] )) self.assertLessEqual( sum(len(encode_diagnostic_envelope(item)) + 1 for item in first), MAX_DIAGNOSTIC_AGGREGATE_BYTES, ) def test_diagnostic_rejects_noncanonical_json_and_malformed_material(self): envelope = self.envelope() payload = encode_diagnostic_envelope(envelope) with self.assertRaises(WorkerContractError): decode_diagnostic_envelope(payload + b' ') malformed = DiagnosticMaterial( encoding=MaterialEncoding.BASE64, head='not-base64', tail=None, original_size=3, stored_size=3, sha256='0' * 64, truncated=False, ) with self.assertRaises(WorkerContractError): self.envelope(http=dataclasses.replace(envelope.http, body=malformed)) with self.assertRaises(WorkerContractError): decode_diagnostic_envelopes_ndjson(payload) if __name__ == '__main__': unittest.main()