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

1166 lines
41 KiB
Python

import base64
import binascii
import hashlib
import hmac
import json
import math
import re
from dataclasses import dataclass
from datetime import datetime, timezone
from enum import Enum
WORKER_EVENT_SCHEMA = 1
DIAGNOSTIC_SCHEMA = 1
WORKER_EVENT_TYPE = 'slot.phase'
PROGRESS_OUTBOX_SCHEMA = 1
PROGRESS_OUTBOX_RELATIVE_PATH = 'control/progress-outbox.json'
MAX_DIAGNOSTIC_BODY_BYTES = 16 * 1024
MAX_DIAGNOSTIC_LOG_BYTES = 32 * 1024
MAX_DIAGNOSTIC_ENVELOPE_BYTES = 64 * 1024
MAX_DIAGNOSTICS_PER_ASSIGNMENT = 32
MAX_DIAGNOSTIC_AGGREGATE_BYTES = 256 * 1024
LEGACY_ERROR_MATERIAL_BYTES = 2 * 1024
DIAGNOSTIC_PROJECTION_VERSION = 1
LEGACY_DIAGNOSTIC_FALLBACK_TIMESTAMP = '1970-01-01T00:00:00.000Z'
_SHA256_RE = re.compile(r'^[0-9a-f]{64}$')
_SCAN_EVENT_ID_RE = re.compile(r'^[0-9a-f]{32,64}$')
_SOURCE_RE = re.compile(r'^[a-z0-9][a-z0-9_.-]{0,63}$')
_CODE_RE = re.compile(r'^[A-Za-z0-9][A-Za-z0-9._:-]{0,255}$')
_UTC_TIMESTAMP_RE = re.compile(
r'^\d{4}-\d{2}-\d{2}T\d{2}:\d{2}:\d{2}(?:\.\d{1,6})?Z$'
)
class WorkerContractError(ValueError):
def __init__(self, category):
self.category = category
super().__init__('worker contract value is invalid')
class WorkerPhase(str, Enum):
IDLE = 'idle'
CLAIMING = 'claiming'
ASSIGNED = 'assigned'
WAITING_PERMIT = 'waiting_permit'
PREPARING = 'preparing'
RESOLVING = 'resolving'
DOWNLOADING = 'downloading'
CLONING = 'cloning'
SCANNING = 'scanning'
FILTERING = 'filtering'
CLEANING = 'cleaning'
BUNDLING = 'bundling'
UPLOADING = 'uploading'
AWAITING_RECEIPT = 'awaiting_receipt'
BACKOFF = 'backoff'
DRAINING = 'draining'
STOPPED = 'stopped'
CANONICAL_WORKER_PHASES = tuple(phase.value for phase in WorkerPhase)
# Same-phase events carry coarse progress. They are accepted by
# validate_phase_transition without obscuring the state-changing edges here.
ALLOWED_PHASE_TRANSITIONS = {
WorkerPhase.IDLE: frozenset((
WorkerPhase.CLAIMING, WorkerPhase.DRAINING, WorkerPhase.STOPPED,
)),
WorkerPhase.CLAIMING: frozenset((
WorkerPhase.ASSIGNED, WorkerPhase.IDLE, WorkerPhase.BACKOFF,
WorkerPhase.DRAINING,
)),
WorkerPhase.ASSIGNED: frozenset((
WorkerPhase.PREPARING, WorkerPhase.WAITING_PERMIT,
WorkerPhase.BACKOFF, WorkerPhase.IDLE,
WorkerPhase.UPLOADING, WorkerPhase.DRAINING,
)),
WorkerPhase.WAITING_PERMIT: frozenset((
WorkerPhase.RESOLVING, WorkerPhase.DOWNLOADING,
WorkerPhase.CLONING, WorkerPhase.SCANNING,
WorkerPhase.UPLOADING, WorkerPhase.IDLE, WorkerPhase.DRAINING,
)),
WorkerPhase.PREPARING: frozenset((
WorkerPhase.WAITING_PERMIT, WorkerPhase.RESOLVING,
WorkerPhase.SCANNING, WorkerPhase.CLEANING,
WorkerPhase.BUNDLING, WorkerPhase.UPLOADING, WorkerPhase.IDLE,
WorkerPhase.DRAINING,
)),
WorkerPhase.RESOLVING: frozenset((
WorkerPhase.DOWNLOADING, WorkerPhase.CLONING, WorkerPhase.SCANNING,
WorkerPhase.FILTERING, WorkerPhase.CLEANING, WorkerPhase.BUNDLING,
WorkerPhase.UPLOADING, WorkerPhase.IDLE, WorkerPhase.WAITING_PERMIT,
WorkerPhase.DRAINING,
)),
WorkerPhase.DOWNLOADING: frozenset((
WorkerPhase.SCANNING, WorkerPhase.FILTERING, WorkerPhase.CLEANING,
WorkerPhase.BUNDLING, WorkerPhase.UPLOADING, WorkerPhase.IDLE,
WorkerPhase.WAITING_PERMIT, WorkerPhase.DRAINING,
)),
WorkerPhase.CLONING: frozenset((
WorkerPhase.SCANNING, WorkerPhase.CLEANING, WorkerPhase.BUNDLING,
WorkerPhase.UPLOADING, WorkerPhase.IDLE, WorkerPhase.WAITING_PERMIT,
WorkerPhase.DRAINING,
)),
WorkerPhase.SCANNING: frozenset((
WorkerPhase.RESOLVING, WorkerPhase.DOWNLOADING, WorkerPhase.CLONING,
WorkerPhase.FILTERING, WorkerPhase.CLEANING, WorkerPhase.BUNDLING,
WorkerPhase.UPLOADING, WorkerPhase.IDLE, WorkerPhase.WAITING_PERMIT,
WorkerPhase.DRAINING,
)),
WorkerPhase.FILTERING: frozenset((
WorkerPhase.CLEANING, WorkerPhase.BUNDLING, WorkerPhase.UPLOADING,
WorkerPhase.IDLE, WorkerPhase.WAITING_PERMIT, WorkerPhase.DRAINING,
)),
WorkerPhase.CLEANING: frozenset((
WorkerPhase.BUNDLING, WorkerPhase.UPLOADING, WorkerPhase.IDLE,
WorkerPhase.WAITING_PERMIT, WorkerPhase.DRAINING,
)),
WorkerPhase.BUNDLING: frozenset((
WorkerPhase.UPLOADING, WorkerPhase.BACKOFF, WorkerPhase.IDLE,
WorkerPhase.WAITING_PERMIT, WorkerPhase.DRAINING,
)),
WorkerPhase.UPLOADING: frozenset((
WorkerPhase.AWAITING_RECEIPT, WorkerPhase.BACKOFF, WorkerPhase.DRAINING,
)),
WorkerPhase.AWAITING_RECEIPT: frozenset((
WorkerPhase.IDLE, WorkerPhase.BACKOFF, WorkerPhase.DRAINING,
)),
WorkerPhase.BACKOFF: frozenset((
WorkerPhase.IDLE, WorkerPhase.CLAIMING, WorkerPhase.UPLOADING,
WorkerPhase.AWAITING_RECEIPT, WorkerPhase.DRAINING,
)),
WorkerPhase.DRAINING: frozenset((WorkerPhase.STOPPED,)),
WorkerPhase.STOPPED: frozenset(),
}
class DiagnosticKind(str, Enum):
PROVIDER_HTTP = 'provider_http'
SCANNER_PROCESS = 'scanner_process'
EXCEPTION = 'exception'
STORAGE = 'storage'
PROTOCOL = 'protocol'
ASSIGNMENT = 'assignment'
class DiagnosticCategory(str, Enum):
AUTHORIZATION = 'authorization'
RATE_LIMIT = 'rate_limit'
NOT_FOUND = 'not_found'
NETWORK = 'network'
TIMEOUT = 'timeout'
PROVIDER = 'provider'
SCANNER = 'scanner'
STORAGE = 'storage'
PROTOCOL = 'protocol'
ASSIGNMENT_EXPIRED = 'assignment_expired'
INTERNAL = 'internal'
class AssignmentOutcome(str, Enum):
ACCEPTED = 'accepted'
PREBUNDLE_FAILED = 'prebundle_failed'
EXPIRED = 'expired'
UNFINISHED = 'unfinished'
class ScanOutcome(str, Enum):
CLEAN = 'clean'
FOUND = 'found'
DEGRADED = 'degraded'
ERROR = 'error'
SKIPPED = 'skipped'
UNAVAILABLE = 'unavailable'
class MaterialEncoding(str, Enum):
TEXT = 'text'
BASE64 = 'base64'
@dataclass(frozen=True, slots=True)
class WorkerEvent:
schema: int
sequence: int
timestamp: str
instance_id: str
slot_id: int
reservation_id: int | None
source: str | None
type: str
phase: WorkerPhase
phase_started_at: str
scan_deadline_at: str | None
assignment_deadline_at: str | None
progress: dict
@dataclass(frozen=True, slots=True)
class DiagnosticMaterial:
encoding: MaterialEncoding
head: str
tail: str | None
original_size: int
stored_size: int
sha256: str
truncated: bool
@dataclass(frozen=True, slots=True)
class DiagnosticHTTPContext:
operation: str
status_code: int
content_type: str | None
request_id: str | None
body: DiagnosticMaterial | None
headers: DiagnosticMaterial | None = None
@dataclass(frozen=True, slots=True)
class DiagnosticProcessContext:
name: str
exit_code: int | None
signal: int | None
timed_out: bool
stdout: DiagnosticMaterial | None
stderr: DiagnosticMaterial | None
@dataclass(frozen=True, slots=True)
class DiagnosticExceptionContext:
type: str
message: str
fingerprint: str
@dataclass(frozen=True, slots=True)
class DiagnosticEnvelope:
schema: int
diagnostic_uid: str
occurrence_id: str
reservation_id: int
scan_event_id: str | None
slot_id: int
source: str
phase: WorkerPhase
kind: DiagnosticKind
category: DiagnosticCategory
code: str
summary: str
retryable: bool
attempt: int
assignment_outcome: AssignmentOutcome | None
scan_outcome: ScanOutcome | None
occurred_at: str
captured_at: str
received_at: str | None
http: DiagnosticHTTPContext | None
process: DiagnosticProcessContext | None
exception: DiagnosticExceptionContext | None
_WORKER_EVENT_FIELDS = frozenset((
'schema', 'sequence', 'timestamp', 'instance_id', 'slot_id',
'reservation_id', 'source', 'type', 'phase', 'phase_started_at',
'scan_deadline_at', 'assignment_deadline_at', 'progress',
))
_MATERIAL_FIELDS = frozenset((
'encoding', 'head', 'tail', 'original_size', 'stored_size', 'sha256',
'truncated',
))
_HTTP_FIELDS = frozenset((
'operation', 'status_code', 'content_type', 'request_id', 'body', 'headers',
))
_LEGACY_HTTP_FIELDS = _HTTP_FIELDS - {'headers'}
_PROCESS_FIELDS = frozenset((
'name', 'exit_code', 'signal', 'timed_out', 'stdout', 'stderr',
))
_EXCEPTION_FIELDS = frozenset(('type', 'message', 'fingerprint'))
_DIAGNOSTIC_FIELDS = frozenset((
'schema', 'diagnostic_uid', 'occurrence_id', 'reservation_id',
'scan_event_id', 'slot_id', 'source', 'phase', 'kind', 'category', 'code',
'summary', 'retryable', 'attempt', 'assignment_outcome', 'scan_outcome',
'occurred_at', 'captured_at', 'received_at', 'http', 'process', 'exception',
))
def _canonical_json(value):
try:
return json.dumps(
value, ensure_ascii=True, sort_keys=True, separators=(',', ':'),
allow_nan=False,
).encode('ascii')
except (TypeError, ValueError, UnicodeError) as exc:
raise WorkerContractError('json') from exc
def _strict_json(payload, *, maximum=None):
if type(payload) is not bytes or not payload:
raise WorkerContractError('json')
if maximum is not None and len(payload) > maximum:
raise WorkerContractError('bounds')
def reject_duplicate(pairs):
result = {}
for key, value in pairs:
if key in result:
raise WorkerContractError('duplicate_field')
result[key] = value
return result
try:
value = json.loads(
payload.decode('utf-8', errors='strict'),
object_pairs_hook=reject_duplicate,
parse_constant=lambda _value: (_ for _ in ()).throw(
WorkerContractError('constant')
),
)
except WorkerContractError:
raise
except (UnicodeDecodeError, json.JSONDecodeError, TypeError, ValueError) as exc:
raise WorkerContractError('json') from exc
if not isinstance(value, dict):
raise WorkerContractError('shape')
if not hmac.compare_digest(_canonical_json(value), payload):
raise WorkerContractError('canonical')
return value
def _enum(value, kind, field):
try:
return kind(value)
except (TypeError, ValueError) as exc:
raise WorkerContractError(field) from exc
def _integer(value, field, *, minimum=None, optional=False):
if optional and value is None:
return None
if type(value) is not int or (minimum is not None and value < minimum):
raise WorkerContractError(field)
return value
def _string(
value, field, *, optional=False, empty=False, maximum=None, pattern=None,
):
if optional and value is None:
return None
if (
not isinstance(value, str)
or '\x00' in value
or (not empty and not value.strip())
or (maximum is not None and len(value) > maximum)
or (pattern is not None and pattern.fullmatch(value) is None)
):
raise WorkerContractError(field)
return value
def validate_utc_timestamp(value, field='timestamp', *, optional=False):
if optional and value is None:
return None
if not isinstance(value, str) or _UTC_TIMESTAMP_RE.fullmatch(value) is None:
raise WorkerContractError(field)
try:
datetime.fromisoformat(value[:-1] + '+00:00')
except ValueError as exc:
raise WorkerContractError(field) from exc
return value
def _timestamp_value(value):
return datetime.fromisoformat(value[:-1] + '+00:00')
def _validate_json_value(value):
if value is None or type(value) in (bool, int, str):
return
if type(value) is float:
if not math.isfinite(value):
raise WorkerContractError('progress')
return
if isinstance(value, list):
for item in value:
_validate_json_value(item)
return
if isinstance(value, dict):
for key, item in value.items():
if not isinstance(key, str):
raise WorkerContractError('progress')
_validate_json_value(item)
return
raise WorkerContractError('progress')
def _worker_event_value(event):
if not isinstance(event, WorkerEvent):
raise WorkerContractError('event_type')
return {
'schema': event.schema,
'sequence': event.sequence,
'timestamp': event.timestamp,
'instance_id': event.instance_id,
'slot_id': event.slot_id,
'reservation_id': event.reservation_id,
'source': event.source,
'type': event.type,
'phase': event.phase.value if isinstance(event.phase, WorkerPhase) else event.phase,
'phase_started_at': event.phase_started_at,
'scan_deadline_at': event.scan_deadline_at,
'assignment_deadline_at': event.assignment_deadline_at,
'progress': event.progress,
}
def _normalize_worker_event(value):
if not isinstance(value, dict) or set(value) != _WORKER_EVENT_FIELDS:
raise WorkerContractError('shape')
if type(value.get('schema')) is not int or value['schema'] != WORKER_EVENT_SCHEMA:
raise WorkerContractError('schema')
timestamp = validate_utc_timestamp(value.get('timestamp'))
phase_started_at = validate_utc_timestamp(value.get('phase_started_at'), 'phase_started_at')
if _timestamp_value(phase_started_at) > _timestamp_value(timestamp):
raise WorkerContractError('phase_started_at')
progress = value.get('progress')
if not isinstance(progress, dict):
raise WorkerContractError('progress')
_validate_json_value(progress)
return WorkerEvent(
schema=WORKER_EVENT_SCHEMA,
sequence=_integer(value.get('sequence'), 'sequence', minimum=1),
timestamp=timestamp,
instance_id=_string(value.get('instance_id'), 'instance_id'),
slot_id=_integer(value.get('slot_id'), 'slot_id', minimum=0),
reservation_id=_integer(
value.get('reservation_id'), 'reservation_id', minimum=1, optional=True,
),
source=_string(value.get('source'), 'source', optional=True),
type=(
value.get('type')
if value.get('type') == WORKER_EVENT_TYPE
else (_ for _ in ()).throw(WorkerContractError('type'))
),
phase=_enum(value.get('phase'), WorkerPhase, 'phase'),
phase_started_at=phase_started_at,
scan_deadline_at=validate_utc_timestamp(
value.get('scan_deadline_at'), 'scan_deadline_at', optional=True,
),
assignment_deadline_at=validate_utc_timestamp(
value.get('assignment_deadline_at'), 'assignment_deadline_at', optional=True,
),
progress=progress,
)
def validate_phase_transition(previous, current, *, allow_same_phase=True):
previous = _enum(previous, WorkerPhase, 'previous_phase')
current = _enum(current, WorkerPhase, 'phase')
if current == previous and allow_same_phase:
return current
if current not in ALLOWED_PHASE_TRANSITIONS[previous]:
raise WorkerContractError('phase_transition')
return current
def validate_worker_event_sequence(events, *, previous_sequence=None):
sequence = _integer(
previous_sequence, 'previous_sequence', minimum=0, optional=True,
)
phases = {}
normalized = []
for event in events:
current = _normalize_worker_event(_worker_event_value(event))
if sequence is not None and current.sequence <= sequence:
raise WorkerContractError('sequence')
key = (current.instance_id, current.slot_id)
if key in phases:
validate_phase_transition(phases[key], current.phase)
phases[key] = current.phase
sequence = current.sequence
normalized.append(current)
return tuple(normalized)
def encode_worker_event(event):
normalized = _normalize_worker_event(_worker_event_value(event))
return _canonical_json(_worker_event_value(normalized))
def decode_worker_event(payload):
return _normalize_worker_event(_strict_json(payload))
def encode_worker_events_ndjson(events):
normalized = validate_worker_event_sequence(events)
return b''.join(encode_worker_event(event) + b'\n' for event in normalized)
def decode_worker_events_ndjson(payload, *, previous_sequence=None):
if type(payload) is not bytes:
raise WorkerContractError('ndjson')
if not payload:
return ()
if not payload.endswith(b'\n') or b'\r' in payload:
raise WorkerContractError('ndjson')
events = tuple(decode_worker_event(line) for line in payload[:-1].split(b'\n'))
return validate_worker_event_sequence(events, previous_sequence=previous_sequence)
def _material_value(material):
if not isinstance(material, DiagnosticMaterial):
raise WorkerContractError('material_type')
return {
'encoding': (
material.encoding.value
if isinstance(material.encoding, MaterialEncoding)
else material.encoding
),
'head': material.head,
'tail': material.tail,
'original_size': material.original_size,
'stored_size': material.stored_size,
'sha256': material.sha256,
'truncated': material.truncated,
}
def _decode_material_segment(value, encoding, field):
if not isinstance(value, str):
raise WorkerContractError(field)
if encoding is MaterialEncoding.TEXT:
return value.encode('utf-8')
try:
decoded = base64.b64decode(value.encode('ascii'), validate=True)
except (UnicodeEncodeError, binascii.Error, ValueError) as exc:
raise WorkerContractError(field) from exc
if base64.b64encode(decoded).decode('ascii') != value:
raise WorkerContractError(field)
return decoded
def _normalize_material(value):
if not isinstance(value, dict) or set(value) != _MATERIAL_FIELDS:
raise WorkerContractError('material_shape')
encoding = _enum(value.get('encoding'), MaterialEncoding, 'material_encoding')
head = _string(value.get('head'), 'material_head', empty=True)
tail = _string(value.get('tail'), 'material_tail', optional=True)
head_bytes = _decode_material_segment(head, encoding, 'material_head')
tail_bytes = (
_decode_material_segment(tail, encoding, 'material_tail')
if tail is not None else b''
)
original_size = _integer(value.get('original_size'), 'original_size', minimum=0)
stored_size = _integer(value.get('stored_size'), 'stored_size', minimum=0)
if stored_size != len(head_bytes) + len(tail_bytes) or original_size < stored_size:
raise WorkerContractError('material_size')
truncated = value.get('truncated')
if type(truncated) is not bool:
raise WorkerContractError('truncated')
if truncated != (original_size > stored_size) or (tail is not None and not truncated):
raise WorkerContractError('truncated')
digest = value.get('sha256')
if not isinstance(digest, str) or _SHA256_RE.fullmatch(digest) is None:
raise WorkerContractError('sha256')
if not truncated and not hmac.compare_digest(
hashlib.sha256(head_bytes).hexdigest(), digest,
):
raise WorkerContractError('sha256')
return DiagnosticMaterial(
encoding=encoding,
head=head,
tail=tail,
original_size=original_size,
stored_size=stored_size,
sha256=digest,
truncated=truncated,
)
def _encode_material_parts(parts):
if any(b'\x00' in part for part in parts):
return MaterialEncoding.BASE64, tuple(
base64.b64encode(part).decode('ascii') for part in parts
)
try:
text = tuple(part.decode('utf-8', errors='strict') for part in parts)
except UnicodeDecodeError:
return MaterialEncoding.BASE64, tuple(
base64.b64encode(part).decode('ascii') for part in parts
)
return MaterialEncoding.TEXT, text
def make_diagnostic_material(payload, *, maximum, head_tail=False):
if type(payload) is not bytes:
raise WorkerContractError('material_type')
if type(maximum) is not int or maximum < 0:
raise WorkerContractError('bounds')
original_size = len(payload)
truncated = original_size > maximum
if not truncated:
parts = (payload,)
elif head_tail and maximum:
head_size = (maximum + 1) // 2
tail_size = maximum - head_size
parts = (payload[:head_size], payload[-tail_size:] if tail_size else b'')
else:
parts = (payload[:maximum],)
encoding, encoded = _encode_material_parts(parts)
tail = encoded[1] if len(encoded) == 2 and encoded[1] else None
stored_size = sum(len(part) for part in parts if part)
return DiagnosticMaterial(
encoding=encoding,
head=encoded[0],
tail=tail,
original_size=original_size,
stored_size=stored_size,
sha256=hashlib.sha256(payload).hexdigest(),
truncated=truncated,
)
def make_diagnostic_material_from_chunks(chunks, *, maximum, head_tail=False):
if type(maximum) is not int or maximum < 0:
raise WorkerContractError('bounds')
digest = hashlib.sha256()
total = 0
complete = bytearray()
head_size = (maximum + 1) // 2 if head_tail else maximum
tail_size = maximum - head_size if head_tail else 0
head = bytearray()
tail = bytearray()
for chunk in chunks:
if type(chunk) is not bytes:
raise WorkerContractError('material_type')
digest.update(chunk)
total += len(chunk)
if len(complete) <= maximum:
complete.extend(chunk)
if len(complete) > maximum:
complete.clear()
if len(head) < head_size:
head.extend(chunk[:head_size - len(head)])
if tail_size:
tail.extend(chunk)
if len(tail) > tail_size:
del tail[:-tail_size]
if total <= maximum:
return make_diagnostic_material(
bytes(complete), maximum=maximum, head_tail=head_tail,
)
parts = (bytes(head), bytes(tail)) if tail_size else (bytes(head),)
encoding, encoded = _encode_material_parts(parts)
return DiagnosticMaterial(
encoding=encoding,
head=encoded[0],
tail=(encoded[1] if len(encoded) == 2 and encoded[1] else None),
original_size=total,
stored_size=sum(len(part) for part in parts),
sha256=digest.hexdigest(),
truncated=True,
)
def make_body_material(payload):
return make_diagnostic_material(
payload, maximum=MAX_DIAGNOSTIC_BODY_BYTES, head_tail=False,
)
def make_log_material(payload, *, maximum=MAX_DIAGNOSTIC_LOG_BYTES):
if type(maximum) is not int or not 0 <= maximum <= MAX_DIAGNOSTIC_LOG_BYTES:
raise WorkerContractError('bounds')
return make_diagnostic_material(payload, maximum=maximum, head_tail=True)
def diagnostic_material_bytes(material):
normalized = _normalize_material(_material_value(material))
head = _decode_material_segment(normalized.head, normalized.encoding, 'material_head')
tail = (
_decode_material_segment(normalized.tail, normalized.encoding, 'material_tail')
if normalized.tail is not None else b''
)
return head + tail
def _http_value(context):
if not isinstance(context, DiagnosticHTTPContext):
raise WorkerContractError('http_type')
value = {
'operation': context.operation,
'status_code': context.status_code,
'content_type': context.content_type,
'request_id': context.request_id,
'body': _material_value(context.body) if context.body is not None else None,
}
if context.headers is not None:
value['headers'] = _material_value(context.headers)
return value
def _normalize_http(value):
if (
not isinstance(value, dict)
or set(value) not in (_HTTP_FIELDS, _LEGACY_HTTP_FIELDS)
):
raise WorkerContractError('http_shape')
status = _integer(value.get('status_code'), 'http_status', minimum=100)
if status > 599:
raise WorkerContractError('http_status')
body = value.get('body')
body = _normalize_material(body) if body is not None else None
headers = value.get('headers')
headers = _normalize_material(headers) if headers is not None else None
if body is not None and (
body.stored_size > MAX_DIAGNOSTIC_BODY_BYTES or body.tail is not None
):
raise WorkerContractError('body_bounds')
if headers is not None and (
headers.stored_size > MAX_DIAGNOSTIC_BODY_BYTES
or headers.tail is not None
):
raise WorkerContractError('header_bounds')
return DiagnosticHTTPContext(
operation=_string(
value.get('operation'), 'http_operation', maximum=128,
pattern=_CODE_RE,
),
status_code=status,
content_type=_string(
value.get('content_type'), 'content_type', optional=True, maximum=256,
),
request_id=_string(
value.get('request_id'), 'request_id', optional=True, maximum=512,
),
body=body,
headers=headers,
)
def _process_value(context):
if not isinstance(context, DiagnosticProcessContext):
raise WorkerContractError('process_type')
return {
'name': context.name,
'exit_code': context.exit_code,
'signal': context.signal,
'timed_out': context.timed_out,
'stdout': _material_value(context.stdout) if context.stdout is not None else None,
'stderr': _material_value(context.stderr) if context.stderr is not None else None,
}
def _normalize_process(value):
if not isinstance(value, dict) or set(value) != _PROCESS_FIELDS:
raise WorkerContractError('process_shape')
stdout = value.get('stdout')
stderr = value.get('stderr')
stdout = _normalize_material(stdout) if stdout is not None else None
stderr = _normalize_material(stderr) if stderr is not None else None
if sum(item.stored_size for item in (stdout, stderr) if item is not None) > MAX_DIAGNOSTIC_LOG_BYTES:
raise WorkerContractError('log_bounds')
timed_out = value.get('timed_out')
if type(timed_out) is not bool:
raise WorkerContractError('timed_out')
return DiagnosticProcessContext(
name=_string(value.get('name'), 'process_name', maximum=256),
exit_code=_integer(value.get('exit_code'), 'exit_code', optional=True),
signal=_integer(value.get('signal'), 'signal', minimum=1, optional=True),
timed_out=timed_out,
stdout=stdout,
stderr=stderr,
)
def _exception_value(context):
if not isinstance(context, DiagnosticExceptionContext):
raise WorkerContractError('exception_type')
return {
'type': context.type,
'message': context.message,
'fingerprint': context.fingerprint,
}
def _normalize_exception(value):
if not isinstance(value, dict) or set(value) != _EXCEPTION_FIELDS:
raise WorkerContractError('exception_shape')
return DiagnosticExceptionContext(
type=_string(value.get('type'), 'exception_type', maximum=512),
message=_string(
value.get('message'), 'exception_message', empty=True, maximum=4096,
),
fingerprint=_string(
value.get('fingerprint'), 'exception_fingerprint', maximum=512,
),
)
def _diagnostic_value(envelope):
if not isinstance(envelope, DiagnosticEnvelope):
raise WorkerContractError('diagnostic_type')
return {
'schema': envelope.schema,
'diagnostic_uid': envelope.diagnostic_uid,
'occurrence_id': envelope.occurrence_id,
'reservation_id': envelope.reservation_id,
'scan_event_id': envelope.scan_event_id,
'slot_id': envelope.slot_id,
'source': envelope.source,
'phase': envelope.phase.value if isinstance(envelope.phase, WorkerPhase) else envelope.phase,
'kind': envelope.kind.value if isinstance(envelope.kind, DiagnosticKind) else envelope.kind,
'category': (
envelope.category.value
if isinstance(envelope.category, DiagnosticCategory)
else envelope.category
),
'code': envelope.code,
'summary': envelope.summary,
'retryable': envelope.retryable,
'attempt': envelope.attempt,
'assignment_outcome': (
envelope.assignment_outcome.value
if isinstance(envelope.assignment_outcome, AssignmentOutcome)
else envelope.assignment_outcome
),
'scan_outcome': (
envelope.scan_outcome.value
if isinstance(envelope.scan_outcome, ScanOutcome)
else envelope.scan_outcome
),
'occurred_at': envelope.occurred_at,
'captured_at': envelope.captured_at,
'received_at': envelope.received_at,
'http': _http_value(envelope.http) if envelope.http is not None else None,
'process': _process_value(envelope.process) if envelope.process is not None else None,
'exception': (
_exception_value(envelope.exception)
if envelope.exception is not None else None
),
}
def _diagnostic_uid(value):
identity = dict(value)
identity.pop('diagnostic_uid', None)
identity.pop('received_at', None)
return hashlib.sha256(_canonical_json(identity)).hexdigest()
def _normalize_diagnostic(value, *, verify_uid=True):
if not isinstance(value, dict) or set(value) != _DIAGNOSTIC_FIELDS:
raise WorkerContractError('shape')
if type(value.get('schema')) is not int or value['schema'] != DIAGNOSTIC_SCHEMA:
raise WorkerContractError('schema')
uid = value.get('diagnostic_uid')
if not isinstance(uid, str) or _SHA256_RE.fullmatch(uid) is None:
raise WorkerContractError('diagnostic_uid')
retryable = value.get('retryable')
if type(retryable) is not bool:
raise WorkerContractError('retryable')
occurred_at = validate_utc_timestamp(value.get('occurred_at'), 'occurred_at')
captured_at = validate_utc_timestamp(value.get('captured_at'), 'captured_at')
if _timestamp_value(captured_at) < _timestamp_value(occurred_at):
raise WorkerContractError('captured_at')
http = value.get('http')
process = value.get('process')
exception = value.get('exception')
http = _normalize_http(http) if http is not None else None
process = _normalize_process(process) if process is not None else None
exception = _normalize_exception(exception) if exception is not None else None
kind = _enum(value.get('kind'), DiagnosticKind, 'kind')
required_context = {
DiagnosticKind.PROVIDER_HTTP: http,
DiagnosticKind.SCANNER_PROCESS: process,
DiagnosticKind.EXCEPTION: exception,
}.get(kind, True)
if required_context is None:
raise WorkerContractError('diagnostic_context')
scan_event_id = _string(
value.get('scan_event_id'), 'scan_event_id', optional=True,
)
if scan_event_id is not None and _SCAN_EVENT_ID_RE.fullmatch(scan_event_id) is None:
raise WorkerContractError('scan_event_id')
normalized = DiagnosticEnvelope(
schema=DIAGNOSTIC_SCHEMA,
diagnostic_uid=uid,
occurrence_id=_string(
value.get('occurrence_id'), 'occurrence_id', maximum=512,
),
reservation_id=_integer(value.get('reservation_id'), 'reservation_id', minimum=1),
scan_event_id=scan_event_id,
slot_id=_integer(value.get('slot_id'), 'slot_id', minimum=0),
source=_string(
value.get('source'), 'source', maximum=64, pattern=_SOURCE_RE,
),
phase=_enum(value.get('phase'), WorkerPhase, 'phase'),
kind=kind,
category=_enum(value.get('category'), DiagnosticCategory, 'category'),
code=_string(
value.get('code'), 'code', maximum=256, pattern=_CODE_RE,
),
summary=_string(value.get('summary'), 'summary'),
retryable=retryable,
attempt=_integer(value.get('attempt'), 'attempt', minimum=1),
assignment_outcome=(
_enum(value.get('assignment_outcome'), AssignmentOutcome, 'assignment_outcome')
if value.get('assignment_outcome') is not None else None
),
scan_outcome=(
_enum(value.get('scan_outcome'), ScanOutcome, 'scan_outcome')
if value.get('scan_outcome') is not None else None
),
occurred_at=occurred_at,
captured_at=captured_at,
received_at=validate_utc_timestamp(
value.get('received_at'), 'received_at', optional=True,
),
http=http,
process=process,
exception=exception,
)
if verify_uid and not hmac.compare_digest(
normalized.diagnostic_uid, _diagnostic_uid(_diagnostic_value(normalized)),
):
raise WorkerContractError('diagnostic_uid')
return normalized
def build_diagnostic_envelope(
*, occurrence_id, reservation_id, scan_event_id, slot_id, source, phase,
kind, category, code, summary, retryable, attempt, assignment_outcome,
scan_outcome, occurred_at, captured_at, received_at=None, http=None,
process=None, exception=None,
):
envelope = DiagnosticEnvelope(
schema=DIAGNOSTIC_SCHEMA,
diagnostic_uid='0' * 64,
occurrence_id=occurrence_id,
reservation_id=reservation_id,
scan_event_id=scan_event_id,
slot_id=slot_id,
source=source,
phase=phase,
kind=kind,
category=category,
code=code,
summary=summary,
retryable=retryable,
attempt=attempt,
assignment_outcome=assignment_outcome,
scan_outcome=scan_outcome,
occurred_at=occurred_at,
captured_at=captured_at,
received_at=received_at,
http=http,
process=process,
exception=exception,
)
normalized = _normalize_diagnostic(_diagnostic_value(envelope), verify_uid=False)
value = _diagnostic_value(normalized)
value['diagnostic_uid'] = _diagnostic_uid(value)
return _normalize_diagnostic(value)
def diagnostic_uid_for(envelope):
normalized = _normalize_diagnostic(_diagnostic_value(envelope), verify_uid=False)
return _diagnostic_uid(_diagnostic_value(normalized))
def encode_diagnostic_envelope(envelope):
normalized = _normalize_diagnostic(_diagnostic_value(envelope))
payload = _canonical_json(_diagnostic_value(normalized))
if len(payload) > MAX_DIAGNOSTIC_ENVELOPE_BYTES:
raise WorkerContractError('envelope_bounds')
return payload
def decode_diagnostic_envelope(payload):
value = _strict_json(payload, maximum=MAX_DIAGNOSTIC_ENVELOPE_BYTES)
return _normalize_diagnostic(value)
def validate_diagnostic_envelopes(envelopes):
normalized = tuple(
_normalize_diagnostic(_diagnostic_value(envelope)) for envelope in envelopes
)
if len(normalized) > MAX_DIAGNOSTICS_PER_ASSIGNMENT:
raise WorkerContractError('diagnostic_count')
total = sum(len(encode_diagnostic_envelope(item)) + 1 for item in normalized)
if total > MAX_DIAGNOSTIC_AGGREGATE_BYTES:
raise WorkerContractError('diagnostic_aggregate')
return normalized
def encode_diagnostic_envelopes_ndjson(envelopes):
normalized = validate_diagnostic_envelopes(envelopes)
return b''.join(encode_diagnostic_envelope(item) + b'\n' for item in normalized)
def decode_diagnostic_envelopes_ndjson(payload):
if type(payload) is not bytes:
raise WorkerContractError('ndjson')
if not payload:
return ()
if (
len(payload) > MAX_DIAGNOSTIC_AGGREGATE_BYTES
or not payload.endswith(b'\n')
or b'\r' in payload
):
raise WorkerContractError('diagnostic_aggregate')
envelopes = tuple(
decode_diagnostic_envelope(line) for line in payload[:-1].split(b'\n')
)
return validate_diagnostic_envelopes(envelopes)
def _legacy_timestamp(value):
try:
parsed = datetime.fromisoformat(str(value).replace('Z', '+00:00'))
except (TypeError, ValueError) as exc:
raise WorkerContractError('timestamp') from exc
if parsed.tzinfo is None:
parsed = parsed.replace(tzinfo=timezone.utc)
return parsed.astimezone(timezone.utc).isoformat(
timespec='milliseconds'
).replace('+00:00', 'Z')
def _legacy_summary(raw):
summary = raw[:256].decode('utf-8', errors='ignore').strip()
return (
'legacy E-frame scan error'
if not summary or '\x00' in summary else summary
)
def build_legacy_error_frame_diagnostics(
*, reservation_id, scan_event_id, slot_id, source, timestamp, errors,
retryable=False, attempt=1,
):
values = [str(error).encode('utf-8') for error in errors]
if not values:
return ()
occurred_at = _legacy_timestamp(
timestamp or LEGACY_DIAGNOSTIC_FALLBACK_TIMESTAMP
)
projected = []
individual = min(len(values), MAX_DIAGNOSTICS_PER_ASSIGNMENT)
if len(values) > MAX_DIAGNOSTICS_PER_ASSIGNMENT:
individual -= 1
for index, raw in enumerate(values[:individual]):
digest = hashlib.sha256(raw).hexdigest()
projected.append(build_diagnostic_envelope(
occurrence_id=f'{scan_event_id}:legacy-e:{index}:{digest}',
reservation_id=reservation_id,
scan_event_id=scan_event_id,
slot_id=slot_id,
source=source,
phase=WorkerPhase.SCANNING,
kind=DiagnosticKind.SCANNER_PROCESS,
category=DiagnosticCategory.SCANNER,
code='legacy.error_frame',
summary=_legacy_summary(raw),
retryable=bool(retryable),
attempt=attempt,
assignment_outcome=AssignmentOutcome.ACCEPTED,
scan_outcome=ScanOutcome.ERROR,
occurred_at=occurred_at,
captured_at=occurred_at,
process=DiagnosticProcessContext(
name='legacy-bundle-error-frame',
exit_code=None,
signal=None,
timed_out=False,
stdout=None,
stderr=make_log_material(
raw, maximum=LEGACY_ERROR_MATERIAL_BYTES,
),
),
exception=DiagnosticExceptionContext(
type='truf.diagnostic.LegacyEFrameProjection',
message=(
'canonical diagnostic projected from the exact persisted '
'legacy E-frame'
),
fingerprint=digest,
),
))
if individual < len(values):
remaining = values[individual:]
material = make_diagnostic_material_from_chunks(
(raw + b'\n' for raw in remaining),
maximum=LEGACY_ERROR_MATERIAL_BYTES,
head_tail=True,
)
projected.append(build_diagnostic_envelope(
occurrence_id=(
f'{scan_event_id}:legacy-e-aggregate:{individual}:'
f'{material.sha256}'
),
reservation_id=reservation_id,
scan_event_id=scan_event_id,
slot_id=slot_id,
source=source,
phase=WorkerPhase.SCANNING,
kind=DiagnosticKind.SCANNER_PROCESS,
category=DiagnosticCategory.SCANNER,
code='legacy.error_frame_aggregate',
summary=f'{len(remaining)} additional legacy E-frame scan errors',
retryable=bool(retryable),
attempt=attempt,
assignment_outcome=AssignmentOutcome.ACCEPTED,
scan_outcome=ScanOutcome.ERROR,
occurred_at=occurred_at,
captured_at=occurred_at,
process=DiagnosticProcessContext(
name='legacy-bundle-error-frame-aggregate',
exit_code=None,
signal=None,
timed_out=False,
stdout=None,
stderr=material,
),
exception=DiagnosticExceptionContext(
type='truf.diagnostic.LegacyEFrameAggregateProjection',
message=(
'bounded canonical aggregate projected from remaining exact '
'persisted legacy E-frames'
),
fingerprint=material.sha256,
),
))
return validate_diagnostic_envelopes(projected)
def ordered_diagnostic_uid_set_sha256(diagnostics):
uids = []
for diagnostic in diagnostics:
uid = (
diagnostic.diagnostic_uid
if isinstance(diagnostic, DiagnosticEnvelope)
else diagnostic.get('diagnostic_uid')
if isinstance(diagnostic, dict) else diagnostic
)
if not isinstance(uid, str) or _SHA256_RE.fullmatch(uid) is None:
raise WorkerContractError('diagnostic_uid')
uids.append(uid)
if len(uids) != len(set(uids)):
raise WorkerContractError('diagnostic_uid')
return hashlib.sha256(_canonical_json(uids)).hexdigest()
worker_event_to_json = encode_worker_event
worker_event_from_json = decode_worker_event
worker_events_to_ndjson = encode_worker_events_ndjson
worker_events_from_ndjson = decode_worker_events_ndjson
diagnostic_envelope_to_json = encode_diagnostic_envelope
diagnostic_envelope_from_json = decode_diagnostic_envelope
diagnostic_envelopes_to_ndjson = encode_diagnostic_envelopes_ndjson
diagnostic_envelopes_from_ndjson = decode_diagnostic_envelopes_ndjson