Files
truf-server/tests/test_worker_contracts.py
T
2026-09-30 20:30:56 +03:00

432 lines
18 KiB
Python

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()