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

2626 lines
112 KiB
Python

import argparse
from contextlib import nullcontext
import hashlib
import http.client
import json
import ntpath
import os
import posixpath
import re
import secrets
import ssl
import subprocess
import sys
import threading
import time
import traceback
from datetime import datetime, timedelta, timezone
from urllib.parse import urlsplit
sys.dont_write_bytecode = True
if not sys.dont_write_bytecode:
raise RuntimeError('remote worker client could not disable bytecode writes')
from owned_process import OwnedProcess
from process_identity import exact_process_identity_state, verify_retained_process
from result_bundle import BundleReservation, ResultBundleReader, bundle_ready_path
from runtime_security import (
MAX_EXTENDED_PRIVATE_JSON_BYTES,
PrivateFileLock,
atomic_write_private_json,
ensure_private_directory,
read_private_json,
reject_reparse_components,
)
import scanner
from scan_execution import (
ScanExecutionError,
WorkerBuildCompatibility,
validate_protocol2_remote_assignment,
)
from worker_package import verify_worker_package
from worker_contracts import (
AssignmentOutcome,
DiagnosticCategory,
DiagnosticExceptionContext,
DiagnosticHTTPContext,
DiagnosticKind,
DiagnosticProcessContext,
MAX_DIAGNOSTIC_BODY_BYTES,
MAX_DIAGNOSTIC_LOG_BYTES,
ScanOutcome,
PROGRESS_OUTBOX_SCHEMA,
WorkerContractError,
WorkerPhase,
build_diagnostic_envelope,
decode_diagnostic_envelope,
encode_diagnostic_envelope,
make_diagnostic_material,
make_log_material,
)
from worker_assignment_runner import (
RunnerProtocolError,
adopt_runner_bundle,
bind_runner_owner,
bind_transferred_runner_owner,
build_runner_input,
cleanup_runner_root,
create_runner_root,
load_generation_terminal,
load_runner_input,
publish_generation_terminal,
publish_start_gate,
read_runner_events,
runner_paths,
runner_root_name,
transfer_runner_to_janitor,
validate_terminal_against_journal,
)
from worker_local_state import WorkerLocalStateError, prepare_progress_outbox_cursor
MAX_API_RESPONSE_BYTES = 64 * 1024 * 1024
MAX_PENDING_BYTES = MAX_EXTENDED_PRIVATE_JSON_BYTES
DIGEST_RE = re.compile(r'^[a-f0-9]{64}$')
CLIENT_STORAGE_FAILURE_DETAIL = 'local worker storage operation failed'
CLIENT_PROCESS_FAILURE_DETAIL = 'local worker execution failed'
SLOT_STATE_SCHEMA = 2
RUNNER_STOP_TIMEOUT_SECONDS = 10.0
TIMEOUT_BUNDLE_PUBLICATION_SECONDS = 30.0
PROGRESS_RETRY_MAX_SECONDS = 30.0
PROGRESS_FINAL_DRAIN_SECONDS = 5.0
PROGRESS_REQUEST_TIMEOUT_SECONDS = 2.0
NO_WORK_REASONS = frozenset((
'empty_queue', 'assignment_cap', 'dispatch_paused', 'capacity',
'compatibility',
))
TERMINAL_REPORT_MAX_BYTES = 16 * 1024
class WorkerClientError(RuntimeError):
pass
class WorkerAssignmentCompatibilityError(WorkerClientError):
pass
class RunnerContainmentPending(WorkerClientError):
pass
class RunnerStageTimeout(RunnerProtocolError):
def __init__(self, phase, scan_started_at, scan_deadline_at, message):
self.phase = WorkerPhase(phase).value
self.scan_started_at = str(scan_started_at)
self.scan_deadline_at = str(scan_deadline_at)
super().__init__(str(message))
class WorkerHTTPError(WorkerClientError):
def __init__(self, status_code, code, message, body=None):
self.status_code = int(status_code)
self.code = str(code or 'request_rejected')
self.body = body
super().__init__(
f'worker API request failed ({self.status_code}, {self.code}): {message}'
)
class WorkerNetworkError(OSError):
pass
def safe_worker_error_summary(error):
if isinstance(error, WorkerHTTPError):
return f'worker API request failed (HTTP {error.status_code})'
if isinstance(error, WorkerNetworkError):
return 'worker network operation failed'
if isinstance(error, WorkerClientError):
return 'worker protocol or state validation failed'
if isinstance(error, OSError):
return 'local I/O operation failed'
return 'worker operation failed'
class WorkerHTTPClient:
def __init__(self, server_url, token, timeout_seconds=120):
parsed = urlsplit(str(server_url or '').rstrip('/'))
if parsed.scheme != 'https' or not parsed.hostname or parsed.username or parsed.password:
raise ValueError('worker server URL must be an HTTPS origin without credentials')
if parsed.query or parsed.fragment or parsed.path not in ('', '/'):
raise ValueError('worker server URL must not contain a path, query, or fragment')
self.host = parsed.hostname
self.port = parsed.port or 443
self.token = str(token or '')
if not 16 <= len(self.token) <= 512:
raise ValueError('worker token length is invalid')
self.timeout_seconds = max(1, min(3600, int(timeout_seconds)))
self.ssl_context = ssl.create_default_context()
self._progress_connections = set()
self._progress_connections_lock = threading.Lock()
def _connection(self, timeout_seconds=None):
timeout = self.timeout_seconds
if timeout_seconds is not None:
timeout = min(timeout, max(0.05, float(timeout_seconds)))
return http.client.HTTPSConnection(
self.host, self.port, timeout=timeout,
context=self.ssl_context,
)
@staticmethod
def _network_call(operation, *args, **kwargs):
try:
return operation(*args, **kwargs)
except (OSError, http.client.HTTPException) as exc:
raise WorkerNetworkError() from exc
@staticmethod
def _decode_response(response):
payload = WorkerHTTPClient._network_call(
response.read, MAX_API_RESPONSE_BYTES + 1,
)
if len(payload) > MAX_API_RESPONSE_BYTES:
raise WorkerClientError('worker API response exceeded its byte bound')
if not payload:
return None
try:
value = json.loads(payload.decode('utf-8', errors='strict'))
except (UnicodeDecodeError, json.JSONDecodeError) as exc:
raise WorkerClientError('worker API returned invalid JSON') from exc
if not isinstance(value, dict):
raise WorkerClientError('worker API returned a non-object response')
return value
@staticmethod
def _rejection(response, value):
error = (value or {}).get('error') or {}
body = (
json.dumps(value, ensure_ascii=True, sort_keys=True, separators=(',', ':')).encode('utf-8')
if value is not None else None
)
return WorkerHTTPError(
response.status, error.get('code'),
error.get('message') or 'request rejected',
body=body,
)
def _json_request(
self, method, path, payload, expected, response_wait_callback=None,
include_no_work_reason=False, timeout_seconds=None,
cancelable_progress=False,
):
body = json.dumps(
payload, ensure_ascii=True, sort_keys=True, separators=(',', ':'),
).encode('utf-8')
connection = self._connection(
timeout_seconds
) if timeout_seconds is not None else self._connection()
if cancelable_progress:
with self._progress_connections_lock:
self._progress_connections.add(connection)
deadline_timer = None
if timeout_seconds is not None:
deadline_timer = threading.Timer(
max(0.01, float(timeout_seconds)), connection.close,
)
deadline_timer.daemon = True
deadline_timer.start()
try:
self._network_call(
connection.request,
method, path, body=body,
headers={
'Authorization': f'Bearer {self.token}',
'Content-Type': 'application/json',
'Content-Length': str(len(body)),
},
)
if response_wait_callback is not None:
response_wait_callback()
response = self._network_call(connection.getresponse)
value = self._decode_response(response)
if response.status not in expected:
raise self._rejection(response, value)
result = (response.status, value, response.getheader('Retry-After'))
if include_no_work_reason:
return (*result, response.getheader('X-Truf-No-Work-Reason'))
return result
finally:
if deadline_timer is not None:
deadline_timer.cancel()
if cancelable_progress:
with self._progress_connections_lock:
self._progress_connections.discard(connection)
self._network_call(connection.close)
def cancel_progress_requests(self):
with self._progress_connections_lock:
connections = tuple(self._progress_connections)
for connection in connections:
try:
connection.close()
except OSError:
pass
def claim(self, request_id, compatibility):
status, value, retry_after, no_work_reason = self._json_request(
'POST', '/api/v1/worker/claim',
{'request_id': request_id, 'build': compatibility},
{200, 201, 204},
include_no_work_reason=True,
)
if status == 204:
try:
retry_after = int(retry_after)
except (TypeError, ValueError, OverflowError) as exc:
raise WorkerClientError('worker API claim retry delay is invalid') from exc
if not 1 <= retry_after <= 300:
raise WorkerClientError('worker API claim retry delay is invalid')
if no_work_reason is not None and no_work_reason not in NO_WORK_REASONS:
raise WorkerClientError('worker API no-work reason is invalid')
return {
'retry_after_seconds': retry_after,
'reason': no_work_reason,
}
if status == 200:
resolution = dict((value or {}).get('resolution') or {})
if (
int(resolution.get('reservation_id') or 0) <= 0
or resolution.get('resolution') not in {
'bundle_accepted', 'prebundle_report', 'expired',
}
or not DIGEST_RE.fullmatch(str(resolution.get('receipt_id') or ''))
or not re.fullmatch(r'[a-f0-9]{32,64}', str(resolution.get('bundle_id') or ''))
or not re.fullmatch(r'[a-f0-9]{32,64}', str(resolution.get('scan_event_id') or ''))
):
raise WorkerClientError('worker API claim resolution is invalid')
return {'claim_resolution': resolution}
assignment = (value or {}).get('assignment')
if not isinstance(assignment, dict):
raise WorkerClientError('worker API claim response has no assignment')
return assignment
def status(self, reservation_id):
connection = self._connection()
try:
self._network_call(
connection.request,
'GET', f'/api/v1/worker/assignments/{int(reservation_id)}',
headers={'Authorization': f'Bearer {self.token}'},
)
response = self._network_call(connection.getresponse)
value = self._decode_response(response)
if response.status != 200:
raise self._rejection(response, value)
return value
finally:
self._network_call(connection.close)
def progress(self, reservation_id, event, *, timeout_seconds=None):
_, value, _ = self._json_request(
'POST', f'/api/v1/worker/assignments/{int(reservation_id)}/progress',
event, {200},
timeout_seconds=(
PROGRESS_REQUEST_TIMEOUT_SECONDS
if timeout_seconds is None else min(
PROGRESS_REQUEST_TIMEOUT_SECONDS, float(timeout_seconds),
)
),
cancelable_progress=True,
)
if (
not isinstance(value, dict)
or value.get('accepted') is not True
or int(value.get('reservation_id') or 0) != int(reservation_id)
or int(value.get('sequence') or 0) != int(event.get('sequence') or 0)
or not isinstance(value.get('received_at'), str)
or type(value.get('replayed')) is not bool
):
raise WorkerClientError('worker API progress acceptance is invalid')
return value
def terminal(
self, reservation_id, report,
response_wait_callback=None,
):
report = dict(report or {})
encoded = json.dumps(
report, ensure_ascii=True, sort_keys=True, separators=(',', ':'),
).encode('ascii')
if len(encoded) > TERMINAL_REPORT_MAX_BYTES:
raise WorkerClientError('pending terminal report exceeds its byte bound')
_, value, _ = self._json_request(
'POST', f'/api/v1/worker/assignments/{int(reservation_id)}/terminal',
report, {200},
response_wait_callback=response_wait_callback,
)
return value
def upload(self, reservation_id, path, response_wait_callback=None):
reject_reparse_components(path)
if not os.path.isfile(path) or os.path.islink(path):
raise WorkerClientError('pending result bundle is not a regular file')
byte_count = os.path.getsize(path)
digest = hashlib.sha256()
with open(path, 'rb', buffering=0) as handle:
for chunk in iter(lambda: handle.read(1024 * 1024), b''):
digest.update(chunk)
payload_sha256 = digest.hexdigest()
connection = self._connection()
try:
self._network_call(
connection.putrequest,
'PUT', f'/api/v1/worker/assignments/{int(reservation_id)}/bundle',
)
self._network_call(
connection.putheader, 'Authorization', f'Bearer {self.token}',
)
self._network_call(
connection.putheader, 'Content-Type', 'application/octet-stream',
)
self._network_call(
connection.putheader, 'Content-Length', str(byte_count),
)
self._network_call(
connection.putheader, 'X-Truf-Payload-SHA256', payload_sha256,
)
self._network_call(connection.endheaders)
with open(path, 'rb', buffering=0) as handle:
for chunk in iter(lambda: handle.read(1024 * 1024), b''):
self._network_call(connection.send, chunk)
if response_wait_callback is not None:
response_wait_callback()
response = self._network_call(connection.getresponse)
value = self._decode_response(response)
if response.status not in (200, 201):
raise self._rejection(response, value)
if str((value or {}).get('payload_sha256') or '') != payload_sha256:
raise WorkerClientError('worker API receipt does not match the uploaded bundle')
return value
finally:
self._network_call(connection.close)
class ProgressOutbox:
def __init__(self, api, state_dir, event_reader):
if not callable(event_reader):
raise TypeError('progress outbox event reader must be callable')
self.api = api
self.event_reader = event_reader
try:
self.path, self.sequence = prepare_progress_outbox_cursor(
state_dir, create=True,
)
except (OSError, ValueError, WorkerLocalStateError) as exc:
raise WorkerClientError(
'persisted progress outbox cursor is invalid or conflicting'
) from exc
def _checkpoint(self, sequence):
sequence = int(sequence)
if sequence <= self.sequence:
raise WorkerClientError('progress outbox cursor did not advance')
atomic_write_private_json(self.path, {
'schema': PROGRESS_OUTBOX_SCHEMA,
'sequence': sequence,
}, max_bytes=64 * 1024)
self.sequence = sequence
def publish_once(self, *, deadline=None):
events = self.event_reader(self.sequence, 128)
if not events:
return False
for event in events:
sequence = int(event.get('sequence') or 0)
if sequence <= self.sequence:
raise WorkerClientError('progress outbox event order is invalid')
reservation_id = event.get('reservation_id')
if reservation_id is not None:
timeout_seconds = PROGRESS_REQUEST_TIMEOUT_SECONDS
if deadline is not None:
timeout_seconds = deadline - time.monotonic()
if timeout_seconds <= 0:
raise TimeoutError('progress publication deadline elapsed')
timeout_seconds = min(
PROGRESS_REQUEST_TIMEOUT_SECONDS, timeout_seconds,
)
try:
self.api.progress(
int(reservation_id), event,
timeout_seconds=timeout_seconds,
)
except WorkerHTTPError as exc:
if not (
exc.status_code == 410 and exc.code == 'progress_stale'
):
raise
self._checkpoint(sequence)
return True
def run(self, stopping):
retry = 0.5
while not stopping.is_set():
try:
worked = self.publish_once(
deadline=time.monotonic() + PROGRESS_REQUEST_TIMEOUT_SECONDS,
)
except Exception:
stopping.wait(retry)
retry = min(PROGRESS_RETRY_MAX_SECONDS, retry * 2)
continue
retry = 0.5
if not worked:
stopping.wait(0.5)
def drain(self, timeout_seconds=PROGRESS_FINAL_DRAIN_SECONDS):
deadline = time.monotonic() + max(0.0, float(timeout_seconds))
retry = 0.05
while time.monotonic() < deadline:
try:
worked = self.publish_once(deadline=deadline)
except Exception:
remaining = deadline - time.monotonic()
if remaining <= 0:
break
time.sleep(min(retry, remaining))
retry = min(PROGRESS_RETRY_MAX_SECONDS, retry * 2)
continue
retry = 0.05
if not worked:
return True
return False
class WorkerSlot:
def __init__(
self, slot_id, api, compatibility, state_dir, bundle_root, *,
work_root=None, package_runtime=None, claim_enabled=True, event_callback=None,
terminal_callback=None, diagnostic_callback=None,
process_factory=OwnedProcess, monotonic=time.monotonic,
):
self.slot_id = int(slot_id)
self.api = api
self.compatibility = WorkerBuildCompatibility.from_mapping(compatibility)
self.state_dir = os.path.abspath(state_dir)
self.state_path = os.path.join(state_dir, f'slot-{self.slot_id}.json')
self.stale_path = os.path.join(state_dir, f'slot-{self.slot_id}-stale.json')
self.bundle_root = bundle_root
self.work_root = ensure_private_directory(
os.path.abspath(work_root or os.path.join(state_dir, 'work')),
reject_reparse=True,
)
self.package_runtime = dict(package_runtime or {})
self.claim_enabled = bool(claim_enabled)
self.retry_after_seconds = None
self.event_callback = event_callback
self.terminal_callback = terminal_callback
self.diagnostic_callback = diagnostic_callback
self._diagnostics = []
self._transport_diagnostics = []
self._event_phase = None
self._process_factory = process_factory
self._monotonic = monotonic
@staticmethod
def _event_timestamp(value):
if not value:
return None
try:
parsed = datetime.fromisoformat(str(value).replace('Z', '+00:00'))
except ValueError as exc:
raise WorkerClientError('worker assignment deadline is invalid') from exc
if parsed.tzinfo is None:
parsed = parsed.replace(tzinfo=timezone.utc)
return parsed.astimezone(timezone.utc).isoformat().replace('+00:00', 'Z')
@staticmethod
def _assignment_details(state):
assignment = dict((state or {}).get('assignment') or {})
reservation = dict(assignment.get('reservation') or {})
deadlines = dict(assignment.get('deadlines') or {})
source = (
reservation.get('source') or assignment.get('source')
or reservation.get('platform') or assignment.get('platform')
)
return {
'reservation_id': int(reservation.get('reservation_id') or 0) or None,
'source': str(source) if source else None,
'scan_deadline_at': WorkerSlot._event_timestamp(
((state or {}).get('runner') or {}).get('scan_deadline_at')
or assignment.get('scan_deadline_at')
or deadlines.get('scan_deadline_at')
or reservation.get('scan_deadline_at')
),
'assignment_deadline_at': WorkerSlot._event_timestamp(
assignment.get('assignment_deadline_at')
or deadlines.get('assignment_deadline_at')
or reservation.get('remote_expires_at')
),
'attempt': max(1, int(reservation.get('attempts') or 1)),
}
@staticmethod
def _runner_identity(value):
if value is None:
return None
if not isinstance(value, dict) or set(value) != {
'host', 'payload', 'job_membership_verified',
} or type(value.get('job_membership_verified')) is not bool:
raise WorkerClientError('persisted runner identity shape is invalid')
normalized = {'job_membership_verified': value['job_membership_verified']}
for name in ('host', 'payload'):
identity = value.get(name)
if not isinstance(identity, dict) or set(identity) != {
'pid', 'creation_time', 'executable',
} or type(identity.get('pid')) is not int or identity['pid'] <= 0 or not all(
isinstance(identity.get(field), str) and identity[field]
for field in ('creation_time', 'executable')
):
raise WorkerClientError('persisted runner process identity is invalid')
normalized[name] = dict(identity)
return normalized
@staticmethod
def _runner_record(value):
if value is None:
return None
fields = {
'generation', 'input_sha256', 'root_name', 'input_ref', 'events_ref',
'start_ref', 'terminal_ref', 'bundle_ref', 'scan_started_at',
'scan_deadline_at', 'watchdog_deadline_at',
'operation', 'timeout_phase', 'attempt', 'status', 'identity', 'last_event',
'terminal_sha256',
}
if not isinstance(value, dict) or set(value) != fields:
raise WorkerClientError('persisted runner record shape is invalid')
generation = str(value.get('generation') or '')
input_sha256 = str(value.get('input_sha256') or '')
if not re.fullmatch(r'[a-f0-9]{32}', generation) or not DIGEST_RE.fullmatch(input_sha256):
raise WorkerClientError('persisted runner generation identity is invalid')
root_name = str(value.get('root_name') or '')
root_match = re.fullmatch(
r'worker-assignment-(0|[1-9][0-9]*)-([1-9][0-9]*)-([a-f0-9]{32})',
root_name,
)
if root_match is None or root_match.group(3) != generation:
raise WorkerClientError('persisted runner root reference is invalid')
expected = {
'input_ref': f'{root_name}/input.json',
'events_ref': f'{root_name}/events.jsonl',
'start_ref': f'{root_name}/start.json',
'terminal_ref': f'{root_name}/terminal.json',
}
if any(value.get(name) != reference for name, reference in expected.items()):
raise WorkerClientError('persisted runner protocol reference is invalid')
bundle_ref = value.get('bundle_ref')
if bundle_ref is not None and (
not isinstance(bundle_ref, str)
or not bundle_ref.startswith(f'{root_name}/bundle/ready/')
or '\\' in bundle_ref or '..' in bundle_ref.split('/')
):
raise WorkerClientError('persisted runner bundle reference is invalid')
started = str(value.get('scan_started_at') or '')
deadline = str(value.get('scan_deadline_at') or '')
normalized_started = WorkerSlot._event_timestamp(started)
normalized_deadline = WorkerSlot._event_timestamp(deadline)
watchdog_deadline = str(value.get('watchdog_deadline_at') or '')
WorkerSlot._event_timestamp(watchdog_deadline)
if datetime.fromisoformat(normalized_deadline.replace('Z', '+00:00')) <= datetime.fromisoformat(normalized_started.replace('Z', '+00:00')):
raise WorkerClientError('persisted runner deadline ordering is invalid')
if value.get('status') not in {
'created', 'running', 'stopping', 'exited', 'timed_out', 'fenced',
}:
raise WorkerClientError('persisted runner status is invalid')
operation = value.get('operation')
timeout_phase = value.get('timeout_phase')
if operation not in {'execute', 'timeout_bundle'}:
raise WorkerClientError('persisted runner operation is invalid')
if operation == 'execute' and timeout_phase is not None:
raise WorkerClientError('persisted scan runner timeout phase is invalid')
if operation == 'timeout_bundle':
try:
timeout_phase = WorkerPhase(timeout_phase).value
except ValueError as exc:
raise WorkerClientError('persisted timeout runner phase is invalid') from exc
if timeout_phase in {
WorkerPhase.IDLE.value, WorkerPhase.CLAIMING.value,
WorkerPhase.UPLOADING.value, WorkerPhase.AWAITING_RECEIPT.value,
WorkerPhase.BACKOFF.value, WorkerPhase.DRAINING.value,
WorkerPhase.STOPPED.value,
}:
raise WorkerClientError('persisted timeout runner phase is outside the scan stage')
if type(value.get('attempt')) is not int or not 1 <= value['attempt'] <= 3:
raise WorkerClientError('persisted runner attempt is invalid')
identity = WorkerSlot._runner_identity(value.get('identity'))
last_event = value.get('last_event')
if last_event is not None:
from worker_assignment_runner import validate_runner_event
try:
last_event = validate_runner_event(
last_event, generation=generation, input_sha256=input_sha256,
)
except RunnerProtocolError as exc:
raise WorkerClientError('persisted runner event is invalid') from exc
terminal_sha256 = value.get('terminal_sha256')
if terminal_sha256 is not None and not DIGEST_RE.fullmatch(str(terminal_sha256)):
raise WorkerClientError('persisted runner terminal reference is invalid')
return {
**value,
'scan_started_at': started,
'scan_deadline_at': deadline,
'watchdog_deadline_at': watchdog_deadline,
'identity': identity,
'last_event': last_event,
}
@staticmethod
def _slot_state(value, *, slot_id=None):
if not isinstance(value, dict):
raise WorkerClientError('persisted worker slot state is invalid')
state = dict(value)
if 'schema' not in state:
phase = state.get('phase')
legacy = {
'claiming': ({'phase', 'request_id'}, None),
'assigned': ({'phase', 'assignment'}, 'runner'),
'bundle_ready': ({'phase', 'assignment'}, 'runner'),
'terminal_pending': ({'phase', 'assignment', 'terminal'}, 'runner'),
'awaiting_resolution': ({'phase', 'assignment', 'stale'}, 'runner'),
}
expected, runner_field = legacy.get(phase, (None, None))
if expected is None or not expected <= set(state) or set(state) - expected - {'bundle'}:
raise WorkerClientError('legacy worker slot state shape is invalid')
state['schema'] = SLOT_STATE_SCHEMA
if runner_field:
state[runner_field] = None
state['retained_work'] = []
if phase == 'bundle_ready' and 'bundle' not in state:
state['bundle'] = {}
if state.get('schema') != SLOT_STATE_SCHEMA:
raise WorkerClientError('persisted worker slot schema is invalid')
phase = state.get('phase')
required = {
'claiming': {'schema', 'phase', 'request_id'},
'assigned': {'schema', 'phase', 'assignment', 'runner', 'retained_work'},
'bundle_ready': {'schema', 'phase', 'assignment', 'runner', 'retained_work', 'bundle'},
'terminal_pending': {'schema', 'phase', 'assignment', 'runner', 'retained_work', 'terminal'},
'awaiting_resolution': {'schema', 'phase', 'assignment', 'runner', 'retained_work', 'stale'},
}.get(phase)
optional = (
{'bundle', 'terminal', 'transport_conflict'}
if phase == 'awaiting_resolution'
else {'bundle', 'transport_conflict'}
if phase == 'terminal_pending'
else {'transport_conflict'} if phase == 'bundle_ready'
else set()
)
if required is None or not required <= set(state) or set(state) - required - optional:
raise WorkerClientError('persisted worker slot state shape is invalid')
if phase == 'claiming':
if not re.fullmatch(r'[a-f0-9]{32}', str(state.get('request_id') or '')):
raise WorkerClientError('persisted claim request identity is invalid')
return state
if not isinstance(state.get('assignment'), dict):
raise WorkerClientError('persisted worker assignment is invalid')
retained_work = state.get('retained_work')
if (
not isinstance(retained_work, list) or len(retained_work) > 16
or any(
not isinstance(item, str)
or re.fullmatch(
r'abandoned/worker-assignment-[0-9]+-[1-9][0-9]*-[a-f0-9]{32}',
item,
) is None
for item in retained_work
)
or len(set(retained_work)) != len(retained_work)
):
raise WorkerClientError('persisted retained runner work is invalid')
state['runner'] = WorkerSlot._runner_record(state.get('runner'))
if state['runner'] is not None:
root_match = re.fullmatch(
r'worker-assignment-(0|[1-9][0-9]*)-([1-9][0-9]*)-([a-f0-9]{32})',
state['runner']['root_name'],
)
reservation_id = int(
(state['assignment'].get('reservation') or {}).get('reservation_id') or 0
)
if (
int(root_match.group(2)) != reservation_id
or (slot_id is not None and int(root_match.group(1)) != int(slot_id))
):
raise WorkerClientError('persisted runner root conflicts with slot authority')
if 'bundle' in state and not isinstance(state.get('bundle'), dict):
raise WorkerClientError('persisted worker bundle state is invalid')
if 'terminal' in state and (
not isinstance(state.get('terminal'), dict)
or set(state['terminal']) not in (
{'failure_code', 'detail'},
{'failure_code', 'detail', 'diagnostics'},
)
or any(type(state['terminal'].get(name)) is not str for name in ('failure_code', 'detail'))
):
raise WorkerClientError('persisted terminal report is invalid')
if 'terminal' in state and 'diagnostics' in state['terminal']:
diagnostics = state['terminal']['diagnostics']
if not isinstance(diagnostics, list):
raise WorkerClientError('persisted terminal diagnostics are invalid')
try:
normalized = [
json.loads(encode_diagnostic_envelope(
decode_diagnostic_envelope(json.dumps(
diagnostic, ensure_ascii=True, sort_keys=True,
separators=(',', ':'),
).encode('ascii'))
).decode('ascii'))
for diagnostic in diagnostics
]
except (TypeError, ValueError, UnicodeError) as exc:
raise WorkerClientError('persisted terminal diagnostics are invalid') from exc
if normalized != diagnostics:
raise WorkerClientError('persisted terminal diagnostics are not canonical')
if 'stale' in state and (
not isinstance(state.get('stale'), dict)
or set(state['stale']) != {'status_code', 'code'}
):
raise WorkerClientError('persisted stale reconciliation is invalid')
if 'transport_conflict' in state:
conflict = state['transport_conflict']
if (
not isinstance(conflict, dict)
or set(conflict) != {'status_code', 'code', 'attempts'}
or conflict.get('status_code') != 409
or not re.fullmatch(r'[a-z0-9_]{1,64}', str(conflict.get('code') or ''))
or type(conflict.get('attempts')) is not int
or not 1 <= conflict['attempts'] <= 2
):
raise WorkerClientError('persisted transport conflict is invalid')
return state
def _emit(
self, phase, state=None, progress=None, *, timestamp=None,
phase_started_at=None,
):
phase = WorkerPhase(phase)
details = self._assignment_details(state)
measured = {'attempt': details.pop('attempt')}
measured.update(dict(progress or {}))
if state is not None and state.get('retained_work'):
measured['retained_work'] = list(state['retained_work'])
if self.event_callback is not None:
self.event_callback({
'slot_id': self.slot_id,
'phase': phase.value,
**details,
'progress': measured,
'timestamp': timestamp,
'phase_started_at': phase_started_at,
})
self._event_phase = phase
def _resume_events(self, state):
if self._event_phase is not None:
return
phase = str((state or {}).get('phase') or '')
if not phase:
self._emit(WorkerPhase.IDLE, progress={'reason': 'startup'})
return
if phase == 'claiming':
self._emit(WorkerPhase.CLAIMING, state, {'recovered': True})
return
runner = dict((state or {}).get('runner') or {})
if phase == 'assigned' and runner:
# Recovery decides whether this generation completed or must be
# fenced before publishing the first event for the new instance.
return
if phase == 'assigned':
self._emit(WorkerPhase.ASSIGNED, state, {'recovered': True})
return
if phase == 'bundle_ready':
self._emit(WorkerPhase.BACKOFF, state, {
'recovered': True, 'reason': 'bundle_upload_recovery',
})
return
self._emit(WorkerPhase.BACKOFF, state, {
'recovered': True, 'reason': 'terminal_reconciliation_recovery',
})
def _finish_idle(self, state, reason):
self._emit(WorkerPhase.IDLE, progress={'reason': reason})
def retire(self, reason):
if self._event_phase != WorkerPhase.DRAINING:
self._emit(WorkerPhase.DRAINING, progress={'reason': str(reason)})
self._emit(WorkerPhase.STOPPED, progress={'reason': str(reason)})
def _notify_terminal(self, state, outcome, receipt):
if self.terminal_callback is None:
return
details = self._assignment_details(state)
reservation = dict((state.get('assignment') or {}).get('reservation') or {})
reservation_id = int(details['reservation_id'] or 0)
receipt_id = str((receipt or {}).get('receipt_id') or '')
history_id = receipt_id or hashlib.sha256(
f'{reservation_id}:{outcome}:{(receipt or {}).get("code", "")}'.encode('utf-8')
).hexdigest()
completed_at = datetime.now(timezone.utc)
started_at = self._event_timestamp(
reservation.get('remote_issued_at') or reservation.get('issued_at')
)
duration_seconds = None
if started_at:
try:
started = datetime.fromisoformat(str(started_at).replace('Z', '+00:00'))
if started.tzinfo is None:
started = started.replace(tzinfo=timezone.utc)
duration_seconds = max(0.0, (completed_at - started.astimezone(timezone.utc)).total_seconds())
except ValueError:
duration_seconds = None
self.terminal_callback({
'history_id': history_id,
'slot_id': self.slot_id,
'reservation_id': reservation_id,
'source': details['source'],
'outcome': str(outcome),
'receipt': dict(receipt or {}),
'started_at': started_at,
'completed_at': completed_at.isoformat().replace('+00:00', 'Z'),
'duration_seconds': duration_seconds,
'first_sequence': None,
'diagnostics': list(self._diagnostics),
})
def _clear_assignment_diagnostics(self):
self._diagnostics.clear()
self._transport_diagnostics.clear()
@staticmethod
def _captured_bytes(error, *names):
for name in names:
value = getattr(error, name, None)
if value is None:
continue
if isinstance(value, bytes):
return value
if isinstance(value, str):
return value.encode('utf-8')
return None
@staticmethod
def _utf8_prefix(value, maximum):
payload = str(value or '').encode('utf-8')
if len(payload) <= maximum:
return payload.decode('utf-8')
return payload[:maximum].decode('utf-8', errors='ignore')
@staticmethod
def _allocate_bytes(desired, budget):
desired = [max(0, int(value)) for value in desired]
budget = min(sum(desired), max(0, int(budget)))
if not desired or not budget or not sum(desired):
return [0 for _value in desired]
total = sum(desired)
values = [min(value, (budget * value) // total) for value in desired]
remaining = budget - sum(values)
order = sorted(
range(len(desired)),
key=lambda index: (
-((budget * desired[index]) % total), index,
),
)
for index in order:
if remaining <= 0:
break
if values[index] < desired[index]:
values[index] += 1
remaining -= 1
return values
@classmethod
def _diagnostic_material_limits(cls, spec, budget):
body = spec.get('body')
stdout = spec.get('stdout')
stderr = spec.get('stderr')
log_desired = cls._allocate_bytes([
min(len(stdout), MAX_DIAGNOSTIC_LOG_BYTES) if stdout is not None else 0,
min(len(stderr), MAX_DIAGNOSTIC_LOG_BYTES) if stderr is not None else 0,
], MAX_DIAGNOSTIC_LOG_BYTES)
desired = [
min(len(body), MAX_DIAGNOSTIC_BODY_BYTES) if body is not None else 0,
*log_desired,
]
return cls._allocate_bytes(desired, budget)
@classmethod
def _build_diagnostic_from_spec(cls, spec, material_budget):
body_limit, stdout_limit, stderr_limit = cls._diagnostic_material_limits(
spec, material_budget,
)
process = None
if (
spec.get('stdout') is not None
or spec.get('stderr') is not None
or spec.get('return_code') is not None
):
process = DiagnosticProcessContext(
name=spec['process_name'],
exit_code=spec['return_code'],
signal=None,
timed_out=spec['timed_out'],
stdout=(
make_log_material(spec['stdout'], maximum=stdout_limit)
if spec.get('stdout') is not None else None
),
stderr=(
make_log_material(spec['stderr'], maximum=stderr_limit)
if spec.get('stderr') is not None else None
),
)
http = None
if spec.get('status_code') is not None:
http = DiagnosticHTTPContext(
operation='worker-api',
status_code=spec['status_code'],
content_type=spec.get('content_type'),
request_id=spec.get('request_id'),
body=(
make_diagnostic_material(spec['body'], maximum=body_limit)
if spec.get('body') is not None else None
),
)
return build_diagnostic_envelope(
occurrence_id=spec['occurrence_id'],
reservation_id=spec['reservation_id'],
scan_event_id=spec['scan_event_id'],
slot_id=spec['slot_id'],
source=spec['source'],
phase=spec['phase'],
kind=spec['kind'],
category=spec['category'],
code=spec['code'],
summary=spec['summary'],
retryable=False,
attempt=spec['attempt'],
assignment_outcome=AssignmentOutcome.PREBUNDLE_FAILED,
scan_outcome=ScanOutcome.UNAVAILABLE,
occurred_at=spec['timestamp'],
captured_at=spec['timestamp'],
http=http,
process=process,
exception=DiagnosticExceptionContext(
type=spec['exception_type'],
message=spec['exception_message'],
fingerprint=spec['fingerprint'],
),
)
def _fitted_terminal_report(self, failure_code, detail):
base = {
'failure_code': str(failure_code),
'detail': str(detail or '')[:1000],
}
if not self._transport_diagnostics:
return base
desired = [
sum(self._diagnostic_material_limits(
spec, MAX_DIAGNOSTIC_BODY_BYTES + MAX_DIAGNOSTIC_LOG_BYTES,
))
for spec in self._transport_diagnostics
]
def candidate(total_budget):
budgets = self._allocate_bytes(desired, total_budget)
value = dict(base)
try:
value['diagnostics'] = [
json.loads(encode_diagnostic_envelope(
self._build_diagnostic_from_spec(spec, budget)
).decode('ascii'))
for spec, budget in zip(self._transport_diagnostics, budgets)
]
except ValueError:
return value, TERMINAL_REPORT_MAX_BYTES + 1
payload = json.dumps(
value, ensure_ascii=True, sort_keys=True, separators=(',', ':'),
).encode('ascii')
return value, len(payload)
low, high = 0, sum(desired)
fitted = None
while low <= high:
middle = (low + high) // 2
value, size = candidate(middle)
if size <= TERMINAL_REPORT_MAX_BYTES:
fitted = value
low = middle + 1
else:
high = middle - 1
if fitted is None:
raise WorkerClientError('canonical terminal report cannot fit its byte bound')
return fitted
def _archive_fitted_terminal_diagnostics(self, report):
if self.diagnostic_callback is None:
return
values = report.get('diagnostics') or []
if len(values) != len(self._transport_diagnostics):
raise WorkerClientError(
'fitted terminal diagnostics lost their local evidence identity'
)
for spec, value in zip(self._transport_diagnostics, values):
envelope = decode_diagnostic_envelope(json.dumps(
value, ensure_ascii=True, sort_keys=True, separators=(',', ':'),
).encode('ascii'))
try:
reference = self.diagnostic_callback(envelope, {
'body': spec.get('body'),
'stdout': spec.get('stdout'),
'stderr': spec.get('stderr'),
})
except Exception:
continue
if reference is not None:
self._diagnostics.append(reference)
def _capture_failure_diagnostic(
self, state, error, code, category, *, operation='assignment_execute',
):
details = self._assignment_details(state)
assignment = dict(state.get('assignment') or {})
reservation = dict(assignment.get('reservation') or {})
body = self._captured_bytes(error, 'body', 'response_body')
stdout = self._captured_bytes(error, 'stdout', 'output')
stderr = self._captured_bytes(error, 'stderr')
return_code = getattr(error, 'returncode', None)
status_code = getattr(error, 'status_code', None)
now = datetime.now(timezone.utc).isoformat().replace('+00:00', 'Z')
if isinstance(error, WorkerContractError):
field = str(error.category)
code = f'worker_contract_invalid.{field}'
category = DiagnosticCategory.PROTOCOL
full_message = f'worker contract field is invalid: {field}'
else:
full_message = str(error)
message = self._utf8_prefix(full_message, 1000)
runner = dict(state.get('runner') or {})
exception_detail = {
'message': full_message,
'operation': str(operation),
'phase': (self._event_phase or WorkerPhase.ASSIGNED).value,
'runner_generation': runner.get('generation'),
'errno': getattr(error, 'errno', None),
'winerror': getattr(error, 'winerror', None),
'filename': getattr(error, 'filename', None),
'filename2': getattr(error, 'filename2', None),
'traceback': ''.join(traceback.format_exception(
type(error), error, error.__traceback__,
)),
}
exception_message = self._utf8_prefix(json.dumps(
exception_detail, ensure_ascii=False, sort_keys=True,
separators=(',', ':'),
), 4000)
fingerprint = hashlib.sha256(
f'{type(error).__module__}.{type(error).__qualname__}:{full_message}'.encode('utf-8')
).hexdigest()
spec = {
'occurrence_id': secrets.token_hex(16),
'reservation_id': int(details['reservation_id'] or 0),
'scan_event_id': str(reservation.get('scan_event_id') or '') or None,
'slot_id': self.slot_id,
'source': str(details['source'] or reservation.get('platform') or 'unknown'),
'phase': self._event_phase or WorkerPhase.ASSIGNED,
'kind': (
DiagnosticKind.SCANNER_PROCESS
if stdout is not None or stderr is not None or type(return_code) is int
else DiagnosticKind.PROVIDER_HTTP
if type(status_code) is int else DiagnosticKind.EXCEPTION
),
'category': category,
'code': str(code),
'summary': message or type(error).__name__,
'attempt': details['attempt'],
'timestamp': now,
'exception_type': f'{type(error).__module__}.{type(error).__qualname__}',
'exception_message': exception_message,
'fingerprint': fingerprint,
'process_name': str(getattr(error, 'process_name', None) or 'worker-operation'),
'return_code': int(return_code) if type(return_code) is int else None,
'timed_out': bool(getattr(error, 'timed_out', False)),
'status_code': int(status_code) if type(status_code) is int else None,
'content_type': getattr(error, 'content_type', None),
'request_id': getattr(error, 'request_id', None),
'body': body,
'stdout': stdout,
'stderr': stderr,
}
self._transport_diagnostics.append(spec)
return None
def _save(self, value):
normalized = self._slot_state(value, slot_id=self.slot_id)
atomic_write_private_json(
self.state_path,
normalized,
max_bytes=MAX_PENDING_BYTES,
)
return normalized
def _load(self):
if not os.path.exists(self.state_path):
return None
loaded = read_private_json(self.state_path, max_bytes=MAX_PENDING_BYTES)
normalized = self._slot_state(loaded, slot_id=self.slot_id)
if normalized != loaded:
atomic_write_private_json(
self.state_path, normalized, max_bytes=MAX_PENDING_BYTES,
)
return normalized
def _remove_state(self):
if os.path.lexists(self.state_path):
reject_reparse_components(self.state_path)
os.remove(self.state_path)
def _resolved(self, state, status):
if not isinstance(status, dict) or not status.get('resolution'):
return False
assignment = state.get('assignment') or {}
reservation = assignment.get('reservation') or {}
reservation_id = int(reservation.get('reservation_id') or 0)
bundle_id = str(reservation.get('bundle_id') or '')
scan_event_id = str(reservation.get('scan_event_id') or '')
resolution = str(status.get('resolution') or '')
if (
int(status.get('reservation_id') or 0) != reservation_id
or str(status.get('bundle_id') or '') != bundle_id
or str(status.get('scan_event_id') or '') != scan_event_id
or resolution not in {
'bundle_accepted', 'prebundle_report', 'expired',
}
or not DIGEST_RE.fullmatch(str(status.get('receipt_id') or ''))
):
raise WorkerClientError('worker API receipt identity is invalid')
phase = str(state.get('phase') or '')
if resolution == 'bundle_accepted':
payload_sha256 = str(status.get('payload_sha256') or '')
if phase not in ('assigned', 'bundle_ready', 'awaiting_resolution') or not DIGEST_RE.fullmatch(
payload_sha256
):
raise WorkerClientError('worker bundle receipt does not match pending state')
path = bundle_ready_path(self.bundle_root, bundle_id)
if os.path.lexists(path):
reject_reparse_components(path)
if not os.path.isfile(path) or os.path.islink(path):
raise WorkerClientError('resolved bundle path is not a regular file')
metadata = ResultBundleReader(
path, max_event_bytes=int(reservation.get('declared_bundle_bytes') or 0),
).validate()
if metadata.header != BundleReservation.from_mapping(reservation).header():
raise WorkerClientError('resolved bundle identity conflicts with pending state')
digest = hashlib.sha256()
with open(path, 'rb', buffering=0) as handle:
for chunk in iter(lambda: handle.read(1024 * 1024), b''):
digest.update(chunk)
if digest.hexdigest() != payload_sha256:
raise WorkerClientError('worker bundle receipt payload digest conflicts')
elif resolution == 'prebundle_report':
terminal = dict(state.get('terminal') or {})
if phase != 'terminal_pending' or set(terminal) not in (
{'failure_code', 'detail'},
{'failure_code', 'detail', 'diagnostics'},
):
raise WorkerClientError('worker terminal receipt does not match pending state')
encoded = json.dumps(
terminal, ensure_ascii=True, sort_keys=True, separators=(',', ':'),
).encode('utf-8')
if (
str(status.get('failure_code') or '') != terminal['failure_code']
or str(status.get('payload_sha256') or '') != hashlib.sha256(encoded).hexdigest()
):
raise WorkerClientError('worker terminal receipt payload conflicts')
if bundle_id:
path = bundle_ready_path(self.bundle_root, bundle_id)
if os.path.lexists(path):
reject_reparse_components(path)
if not os.path.isfile(path) or os.path.islink(path):
raise WorkerClientError('resolved bundle path is not a regular file')
os.remove(path)
self._notify_terminal(state, resolution, status)
self._clear_assignment_diagnostics()
self._remove_state()
self._finish_idle(state, 'terminal_reconciled')
return True
def _cleanup_recorded_stale(self, state):
if not os.path.exists(self.stale_path):
return False
record = read_private_json(self.stale_path, max_bytes=MAX_PENDING_BYTES)
assignment = dict(state.get('assignment') or {})
reservation = dict(assignment.get('reservation') or {})
if not isinstance(record, dict) or set(record) != {
'schema', 'outcome', 'reservation_id', 'bundle_id', 'scan_event_id',
'status_code', 'code', 'recorded_at',
} or (
record.get('schema') != 1
or record.get('outcome') != 'discarded_stale'
or int(record.get('reservation_id') or 0) != int(reservation.get('reservation_id') or 0)
or str(record.get('bundle_id') or '') != str(reservation.get('bundle_id') or '')
or str(record.get('scan_event_id') or '') != str(reservation.get('scan_event_id') or '')
):
return False
bundle_id = str(record['bundle_id'])
if bundle_id:
path = bundle_ready_path(self.bundle_root, bundle_id)
if os.path.lexists(path):
reject_reparse_components(path)
if not os.path.isfile(path) or os.path.islink(path):
raise WorkerClientError('stale bundle path is not a regular file')
os.remove(path)
self._notify_terminal(state, 'discarded_stale', record)
self._clear_assignment_diagnostics()
self._remove_state()
self._finish_idle(state, 'stale_reconciled')
return True
def _mark_stale(self, state, error):
assignment = dict(state.get('assignment') or {})
reservation = dict(assignment.get('reservation') or {})
code = str(error.code or '')
if not re.fullmatch(r'[a-z0-9_]{1,64}', code):
code = 'request_rejected'
record = {
'schema': 1,
'outcome': 'discarded_stale',
'reservation_id': int(reservation.get('reservation_id') or 0),
'bundle_id': str(reservation.get('bundle_id') or ''),
'scan_event_id': str(reservation.get('scan_event_id') or ''),
'status_code': int(error.status_code),
'code': code,
'recorded_at': datetime.now(timezone.utc).isoformat(),
}
if record['reservation_id'] <= 0 or not record['bundle_id'] or not record['scan_event_id']:
raise WorkerClientError('stale assignment identity is invalid')
atomic_write_private_json(
self.stale_path, record, max_bytes=MAX_PENDING_BYTES,
)
if not self._cleanup_recorded_stale(state):
raise WorkerClientError('stale outcome could not be reconciled')
return True
def _reconcile_http_failure(self, state, error):
if error.status_code not in (404, 409, 410):
raise error
try:
status = self.api.status(
int((state.get('assignment') or {}).get('reservation', {}).get('reservation_id') or 0)
)
except WorkerHTTPError as status_error:
if status_error.status_code in (404, 409, 410):
return self._mark_stale(state, status_error)
raise
if self._resolved(state, status):
return True
if error.status_code == 409 and str(status.get('state') or '') == 'scanning':
previous = dict(state.get('transport_conflict') or {})
attempts = (
int(previous.get('attempts') or 0) + 1
if previous.get('code') == str(error.code or '') else 1
)
if attempts == 1:
state['transport_conflict'] = {
'status_code': 409,
'code': str(error.code or 'request_rejected'),
'attempts': 1,
}
self._save(state)
return False
state.pop('transport_conflict', None)
if state.get('phase') == 'bundle_ready':
self._capture_failure_diagnostic(
state, error, 'client_result_conflict',
DiagnosticCategory.PROTOCOL,
operation='bundle_upload_reconciliation',
)
reservation_id = int(
(state.get('assignment') or {}).get(
'reservation', {},
).get('reservation_id') or 0
)
return self._terminal(
state, reservation_id, 'client_process_failed',
CLIENT_PROCESS_FAILURE_DETAIL,
)
state['phase'] = 'awaiting_resolution'
state['stale'] = {
'status_code': error.status_code,
'code': str(error.code or 'request_rejected'),
}
self._save(state)
return False
if error.status_code == 410:
state['phase'] = 'awaiting_resolution'
state['stale'] = {'status_code': error.status_code, 'code': error.code}
self._save(state)
return False
return self._mark_stale(state, error)
def _adopt_ready_bundle(self, state, *, persist=True):
assignment = dict(state['assignment'])
reservation = BundleReservation.from_mapping(assignment['reservation'])
path = bundle_ready_path(self.bundle_root, reservation.bundle_id)
if not os.path.lexists(path):
return None
reject_reparse_components(path)
if not os.path.isfile(path) or os.path.islink(path):
raise WorkerClientError('pending result bundle is not a regular file')
metadata = ResultBundleReader(
path, max_event_bytes=reservation.declared_bytes,
).validate()
if metadata.header != reservation.header():
raise WorkerClientError('pending result bundle identity conflicts with its assignment')
runner = state.get('runner')
if runner is not None:
terminal = self._load_runner_terminal(state)
if (
terminal is None or terminal['decision'] != 'completed'
or terminal['outcome']['status'] != 'succeeded'
):
closed = self._fence_recovered_runner(
state, 'ready_bundle_without_completion',
)
if closed == 'unknown':
raise RunnerContainmentPending(
'ready bundle runner containment remains live',
)
raise RunnerProtocolError(
'canonical ready bundle lacks a valid completed generation',
)
digest = hashlib.sha256()
with open(path, 'rb', buffering=0) as handle:
for block in iter(lambda: handle.read(1024 * 1024), b''):
digest.update(block)
if digest.hexdigest() != terminal['outcome']['bundle']['payload_sha256']:
raise RunnerProtocolError(
'canonical ready bundle conflicts with runner terminal payload',
)
if self._stop_persisted_runner(runner) != 'dead':
raise RunnerContainmentPending(
'completed ready-bundle runner containment remains live',
)
self._retain_runner_work(state, runner)
state['runner'] = None
state['phase'] = 'bundle_ready'
state['bundle'] = metadata.as_dict()
runner = state.get('runner')
if runner is not None:
reference = f"abandoned/{runner['root_name']}"
destination = os.path.join(self.work_root, *reference.split('/'))
if os.path.isdir(destination) and not os.path.exists(
self._runner_paths(runner)['root']
):
if reference not in state.setdefault('retained_work', []):
state['retained_work'].append(reference)
state['runner'] = None
if persist:
self._save(state)
return state
@staticmethod
def _deadline_active(value):
text = str(value or '').strip()
if not text:
return False
try:
parsed = datetime.fromisoformat(text.replace('Z', '+00:00'))
except ValueError as exc:
raise WorkerClientError('worker assignment deadline is invalid') from exc
if parsed.tzinfo is None:
parsed = parsed.replace(tzinfo=timezone.utc)
return parsed.astimezone(timezone.utc) > datetime.now(timezone.utc)
@staticmethod
def _identity_value(value):
return {
name: value[name]
for name in ('pid', 'creation_time', 'executable')
}
def _new_runner(
self, state, *, scan_started_at=None, scan_deadline_at=None,
watchdog_deadline_at=None, operation='execute', timeout_phase=None,
attempt=1,
):
assignment = dict(state['assignment'])
if set(self.package_runtime) != {
'code_manifest', 'code_manifest_sha256', 'trufflehog_path',
'git_path', 'detector_policy_path', 'capabilities', 'bootstrap_path',
}:
raise WorkerClientError('verified worker package runtime is unavailable')
packaged_capabilities = self.package_runtime['capabilities']
if not isinstance(packaged_capabilities, tuple) or not packaged_capabilities:
raise WorkerClientError('verified worker package capabilities are unavailable')
try:
validate_protocol2_remote_assignment(
assignment, self.compatibility, packaged_capabilities,
)
except (ScanExecutionError, TypeError, ValueError) as exc:
raise WorkerAssignmentCompatibilityError(
'worker assignment is incompatible with this package'
) from exc
reservation = dict(assignment['reservation'])
deadlines = dict(assignment['deadlines'])
if scan_started_at is None or scan_deadline_at is None:
started = datetime.now(timezone.utc)
assignment_deadline = self._event_timestamp(
deadlines.get('assignment_deadline_at') or reservation.get('remote_expires_at')
)
assignment_deadline_value = datetime.fromisoformat(
assignment_deadline.replace('Z', '+00:00'),
)
scan_deadline_value = min(
started + timedelta(seconds=int(deadlines['target_scan_timeout_seconds'])),
assignment_deadline_value,
)
if scan_deadline_value <= started:
raise WorkerClientError('worker assignment deadline has passed')
scan_started_at = started.isoformat(timespec='milliseconds').replace('+00:00', 'Z')
scan_deadline_at = scan_deadline_value.isoformat(timespec='milliseconds').replace('+00:00', 'Z')
else:
self._event_timestamp(scan_started_at)
self._event_timestamp(scan_deadline_at)
scan_started_at = str(scan_started_at)
scan_deadline_at = str(scan_deadline_at)
watchdog_deadline_at = str(watchdog_deadline_at or scan_deadline_at)
self._event_timestamp(watchdog_deadline_at)
generation = secrets.token_hex(16)
root_name = runner_root_name(
self.slot_id, reservation['reservation_id'], generation,
)
runner_input = build_runner_input(
assignment,
generation=generation,
slot_id=self.slot_id,
scan_started_at=scan_started_at,
scan_deadline_at=scan_deadline_at,
watchdog_deadline_at=watchdog_deadline_at,
operation=operation,
timeout_phase=timeout_phase,
)
created = threading.Event()
cancelled = threading.Event()
creation = {}
def create_protocol_root():
try:
creation['value'] = create_runner_root(
self.work_root, root_name, runner_input,
)
if cancelled.is_set():
cleanup_runner_root(self.work_root, root_name)
except BaseException as exc:
creation['error'] = exc
finally:
created.set()
threading.Thread(
target=create_protocol_root,
name=f'worker-runner-protocol-{self.slot_id}', daemon=True,
).start()
watchdog_deadline = datetime.fromisoformat(
watchdog_deadline_at.replace('Z', '+00:00'),
)
remaining = max(
0.0, (watchdog_deadline - datetime.now(timezone.utc)).total_seconds(),
)
if not created.wait(remaining):
cancelled.set()
raise RunnerStageTimeout(
WorkerPhase.PREPARING, scan_started_at, scan_deadline_at,
'runner protocol initialization exceeded its absolute deadline',
)
if 'error' in creation:
raise creation['error']
_, input_sha256 = creation['value']
state['runner'] = {
'generation': generation,
'input_sha256': input_sha256,
'root_name': root_name,
'input_ref': f'{root_name}/input.json',
'events_ref': f'{root_name}/events.jsonl',
'start_ref': f'{root_name}/start.json',
'terminal_ref': f'{root_name}/terminal.json',
'bundle_ref': None,
'scan_started_at': scan_started_at,
'scan_deadline_at': scan_deadline_at,
'watchdog_deadline_at': watchdog_deadline_at,
'operation': operation,
'timeout_phase': timeout_phase,
'attempt': int(attempt),
'status': 'created',
'identity': None,
'last_event': None,
'terminal_sha256': None,
}
self._save(state)
return state['runner']
def _runner_paths(self, runner):
return runner_paths(self.work_root, runner['root_name'])
def _drain_runner_events(self, state, *, emit=True):
runner = state['runner']
paths = self._runner_paths(runner)
after = int((runner.get('last_event') or {}).get('sequence') or 0)
events = read_runner_events(
paths['events'], generation=runner['generation'],
input_sha256=runner['input_sha256'], operation=runner['operation'],
after_sequence=after,
)
for event in events:
if emit:
progress = dict(event['progress'])
progress.update({
'runner_generation': runner['generation'],
'runner_sequence': event['sequence'],
})
self._emit(
event['phase'], state, progress,
timestamp=event['timestamp'],
phase_started_at=event['phase_started_at'],
)
runner['last_event'] = event
self._save(state)
return events
@staticmethod
def _stop_exact_identity(identity, timeout):
state = exact_process_identity_state(
identity['pid'], identity['creation_time'], identity['executable'],
)
if state in {'dead', 'reused'}:
return True
if state != 'alive':
return False
try:
with verify_retained_process(
identity['pid'], identity['creation_time'], identity['executable'],
terminate=True,
) as process:
process.terminate()
return process.wait(timeout)
except OSError:
return exact_process_identity_state(
identity['pid'], identity['creation_time'], identity['executable'],
) in {'dead', 'reused'}
def _stop_persisted_runner(self, runner):
identity = runner.get('identity')
if identity is None:
return 'dead'
host_dead = self._stop_exact_identity(
identity['host'], RUNNER_STOP_TIMEOUT_SECONDS / 2,
)
payload_dead = self._stop_exact_identity(
identity['payload'], RUNNER_STOP_TIMEOUT_SECONDS / 2,
)
return 'dead' if host_dead and payload_dead else 'unknown'
def _fence_recovered_runner(self, state, reason):
runner = state['runner']
paths = self._runner_paths(runner)
try:
terminal, _won = publish_generation_terminal(
paths['terminal'], generation=runner['generation'],
input_sha256=runner['input_sha256'], decision='fenced',
reason=str(reason),
)
if terminal['decision'] == 'completed':
return 'completed'
except RunnerProtocolError:
# A malformed complete terminal record is itself a closed but
# unusable generation. Exact containment still has to be stopped.
pass
runner['status'] = 'fenced'
self._save(state)
return self._stop_persisted_runner(runner)
def _runner_command(self, root):
bootstrap = os.path.abspath(self.package_runtime['bootstrap_path'])
if not os.path.isfile(bootstrap) or os.path.islink(bootstrap):
raise WorkerClientError('verified worker bootstrap is unavailable')
return [
sys._base_executable or sys.executable,
'-I', '-S', '-B', bootstrap, '--',
'_assignment_runner', '--root', root,
]
def _load_runner_terminal(self, state):
runner = state['runner']
paths = self._runner_paths(runner)
if not os.path.exists(paths['terminal']):
return None
runner_input, input_digest = load_runner_input(paths['input'])
if input_digest != runner['input_sha256']:
raise RunnerProtocolError('runner input hash conflicts with slot authority')
terminal, digest = load_generation_terminal(
paths['terminal'], generation=runner['generation'],
input_sha256=runner['input_sha256'],
)
events = read_runner_events(
paths['events'], generation=runner['generation'],
input_sha256=runner['input_sha256'], operation=runner['operation'],
)
terminal = validate_terminal_against_journal(
terminal, events, runner_input,
)
if terminal['decision'] != 'completed':
runner['terminal_sha256'] = digest
runner['status'] = (
'timed_out' if terminal['decision'] == 'timed_out' else 'fenced'
)
self._save(state)
return terminal
if runner['status'] in {'stopping', 'timed_out', 'fenced'}:
raise RunnerProtocolError('completed output belongs to a closed runner generation')
deadline_name = (
'scan_deadline_at' if runner['operation'] == 'execute'
else 'watchdog_deadline_at'
)
decided = datetime.fromisoformat(
terminal['decided_at'].replace('Z', '+00:00'),
)
deadline = datetime.fromisoformat(
runner[deadline_name].replace('Z', '+00:00'),
)
if decided > deadline:
raise RunnerProtocolError('runner completed after its absolute deadline')
outcome = terminal['outcome']
assignment = state['assignment']
reservation = assignment['reservation']
identity = outcome['identity']
if identity != {
'slot_id': self.slot_id,
'reservation_id': int(reservation['reservation_id']),
'bundle_id': str(reservation['bundle_id']),
'scan_event_id': str(reservation['scan_event_id']),
'execution_snapshot_sha256': str(assignment['execution_snapshot_sha256']),
}:
raise RunnerProtocolError('runner outcome identity conflicts with slot authority')
self._drain_runner_events(state)
last_sequence = int((runner.get('last_event') or {}).get('sequence') or 0)
if outcome['last_event_sequence'] != last_sequence:
raise RunnerProtocolError('runner outcome event tail is incomplete or inconsistent')
runner['terminal_sha256'] = digest
runner['status'] = 'exited'
if outcome['bundle'] is not None:
runner['bundle_ref'] = (
f"{runner['root_name']}/bundle/"
f"{outcome['bundle']['ready_relative_path']}"
)
self._save(state)
return terminal
def _adopt_runner_outcome(self, state, outcome):
if outcome['status'] != 'succeeded':
self._retain_runner_work(state, state['runner'])
state['runner'] = None
self._save(state)
error = outcome['error']
raise RunnerProtocolError(
f"assignment runner failed in {outcome.get('final_phase') or 'startup'}: "
f"{error['code']}"
)
metadata = adopt_runner_bundle(
self.work_root, state['runner']['root_name'], outcome,
state['assignment'], self.bundle_root,
)
state['phase'] = 'bundle_ready'
state['bundle'] = {
**outcome['bundle']['commit'],
'canonical_scan_event_hash': metadata.scan_event_hash,
}
self._retain_runner_work(state, state['runner'])
state['runner'] = None
self._save(state)
return state
def _retain_runner_work(self, state, runner):
try:
reference = transfer_runner_to_janitor(
self.work_root, runner['root_name'], runner['generation'],
)
except (OSError, RunnerProtocolError) as exc:
raise RunnerContainmentPending(
'runner work has not reached durable janitor ownership',
) from exc
retained = state.setdefault('retained_work', [])
if reference not in retained:
retained.append(reference)
return reference
def _retry_closed_runner(self, state, reason):
runner = state['runner']
if runner['operation'] == 'timeout_bundle':
if runner['attempt'] >= 3:
self._retain_runner_work(state, runner)
state['runner'] = None
self._save(state)
raise RunnerProtocolError(
'timeout bundle runner exhausted its bounded generation retries',
)
return self._timeout_bundle(state)
remaining = (
datetime.fromisoformat(
runner['scan_deadline_at'].replace('Z', '+00:00'),
) - datetime.now(timezone.utc)
).total_seconds()
if remaining < 1:
runner['status'] = 'timed_out'
self._save(state)
return self._timeout_bundle(state)
self._retain_runner_work(state, runner)
if runner['attempt'] >= 3:
state['runner'] = None
self._save(state)
raise RunnerProtocolError('assignment runner exhausted its bounded generation retries')
scan_started_at = runner['scan_started_at']
scan_deadline_at = runner['scan_deadline_at']
attempt = runner['attempt'] + 1
if self._event_phase is not None:
self._emit(self._event_phase, state, {
'runner_retry_reason': str(reason),
'next_runner_attempt': attempt,
})
state['runner'] = None
self._save(state)
try:
self._new_runner(
state,
scan_started_at=scan_started_at,
scan_deadline_at=scan_deadline_at,
watchdog_deadline_at=scan_deadline_at,
operation='execute', attempt=attempt,
)
except RunnerStageTimeout as exc:
return self._start_timeout_bundle(
state,
scan_started_at=exc.scan_started_at,
scan_deadline_at=exc.scan_deadline_at,
final_phase=exc.phase,
)
return self._launch_runner(state)
def _timeout_bundle(self, state):
runner = state['runner']
final_phase = str(
runner.get('timeout_phase')
or (runner.get('last_event') or {}).get('phase')
or 'preparing'
)
scan_started_at = runner['scan_started_at']
scan_deadline_at = runner['scan_deadline_at']
timeout_attempt = (
runner['attempt'] + 1
if runner['operation'] == 'timeout_bundle' else 1
)
self._retain_runner_work(state, runner)
state['runner'] = None
self._save(state)
return self._start_timeout_bundle(
state,
scan_started_at=scan_started_at,
scan_deadline_at=scan_deadline_at,
final_phase=final_phase,
attempt=timeout_attempt,
)
def _start_timeout_bundle(
self, state, *, scan_started_at, scan_deadline_at,
final_phase, attempt=1,
):
assignment = state['assignment']
assignment_deadline = datetime.fromisoformat(
self._event_timestamp(
assignment['deadlines']['assignment_deadline_at'],
).replace('Z', '+00:00'),
)
remaining = (assignment_deadline - datetime.now(timezone.utc)).total_seconds()
if remaining <= 0:
raise RunnerProtocolError('assignment authority expired before timeout publication')
watchdog_deadline_at = (
datetime.now(timezone.utc) + timedelta(seconds=min(
TIMEOUT_BUNDLE_PUBLICATION_SECONDS, remaining,
))
).isoformat(timespec='milliseconds').replace('+00:00', 'Z')
try:
self._new_runner(
state,
scan_started_at=scan_started_at,
scan_deadline_at=scan_deadline_at,
watchdog_deadline_at=watchdog_deadline_at,
operation='timeout_bundle',
timeout_phase=final_phase,
attempt=attempt,
)
except RunnerStageTimeout as exc:
raise RunnerContainmentPending(
'timeout result authority persistence remains pending',
) from exc
return self._launch_runner(state, allow_timeout_fallback=False)
def _launch_process_bounded(self, runner, paths, remaining):
completed = threading.Event()
cancelled = threading.Event()
holder = {}
def launch():
try:
process = self._process_factory(
self._runner_command(paths['root']),
stdin=subprocess.DEVNULL,
stdout=subprocess.DEVNULL,
stderr=subprocess.DEVNULL,
close_fds=True,
creationflags=(
subprocess.CREATE_NO_WINDOW if os.name == 'nt' else 0
),
_startup_timeout=max(0.01, min(15.0, remaining)),
)
if cancelled.is_set():
process.kill()
try:
process.wait(timeout=RUNNER_STOP_TIMEOUT_SECONDS)
except subprocess.TimeoutExpired:
pass
else:
try:
transferred = bind_transferred_runner_owner(
self.work_root, runner['root_name'],
self._identity_value(process.payload_identity),
)
if not transferred:
bind_runner_owner(
self.work_root, runner['root_name'],
self._identity_value(process.payload_identity),
)
except (OSError, RunnerProtocolError):
pass
else:
holder['process'] = process
except BaseException as exc:
holder['error'] = exc
finally:
completed.set()
self._runner_launch_cleanup = completed
threading.Thread(
target=launch, name=f"worker-runner-launch-{self.slot_id}", daemon=True,
).start()
if not completed.wait(max(0.0, remaining)):
cancelled.set()
publish_generation_terminal(
paths['terminal'], generation=runner['generation'],
input_sha256=runner['input_sha256'], decision='timed_out',
reason='startup_deadline',
)
return None
if 'error' in holder:
raise holder['error']
return holder.get('process')
def _launch_runner(self, state, *, allow_timeout_fallback=True):
runner = state['runner']
paths = self._runner_paths(runner)
watchdog_deadline = datetime.fromisoformat(
runner['watchdog_deadline_at'].replace('Z', '+00:00'),
)
remaining = (watchdog_deadline - datetime.now(timezone.utc)).total_seconds()
if remaining <= 0:
publish_generation_terminal(
paths['terminal'], generation=runner['generation'],
input_sha256=runner['input_sha256'], decision='timed_out',
reason='startup_deadline',
)
runner['status'] = 'timed_out'
self._save(state)
if allow_timeout_fallback and runner['operation'] == 'execute':
return self._timeout_bundle(state)
raise RunnerProtocolError('timeout result bundle publication exceeded its watchdog')
try:
process = self._launch_process_bounded(runner, paths, remaining)
except BaseException:
runner['status'] = 'fenced'
self._save(state)
try:
cleanup_runner_root(self.work_root, runner['root_name'])
except OSError:
pass
raise
if process is None:
runner['status'] = 'timed_out'
self._save(state)
if allow_timeout_fallback and runner['operation'] == 'execute':
return self._timeout_bundle(state)
raise RunnerProtocolError('timeout result bundle publication exceeded its watchdog')
watchdog_stop = threading.Event()
watchdog_expired = threading.Event()
def watchdog():
wait = max(
0.0,
(watchdog_deadline - datetime.now(timezone.utc)).total_seconds(),
)
if watchdog_stop.wait(wait):
return
watchdog_expired.set()
process.kill()
try:
publish_generation_terminal(
paths['terminal'], generation=runner['generation'],
input_sha256=runner['input_sha256'], decision='timed_out',
reason=(
'scan_stage_deadline' if runner['operation'] == 'execute'
else 'timeout_bundle_deadline'
),
)
except (OSError, RunnerProtocolError):
pass
watchdog_thread = threading.Thread(
target=watchdog, name=f"worker-runner-watchdog-{self.slot_id}", daemon=True,
)
watchdog_thread.start()
try:
runner['identity'] = {
'host': self._identity_value(process.host_identity),
'payload': self._identity_value(process.payload_identity),
'job_membership_verified': bool(process.job_membership_verified),
}
if not runner['identity']['job_membership_verified']:
raise RunnerProtocolError('assignment runner containment is unverified')
runner['status'] = 'running'
self._save(state)
if watchdog_expired.is_set():
raise RunnerContainmentPending('runner deadline crossed during identity persistence')
bind_runner_owner(
self.work_root, runner['root_name'], runner['identity']['payload'],
)
if watchdog_expired.is_set():
raise RunnerContainmentPending('runner deadline crossed during owner persistence')
publish_start_gate(
paths['start'], generation=runner['generation'],
input_sha256=runner['input_sha256'],
host=runner['identity']['host'], payload=runner['identity']['payload'],
)
while process.poll() is None:
if watchdog_expired.is_set():
break
self._drain_runner_events(state)
time.sleep(0.1)
if watchdog_expired.is_set() and process.poll() is None:
try:
process.wait(timeout=RUNNER_STOP_TIMEOUT_SECONDS)
except subprocess.TimeoutExpired as exc:
runner['status'] = 'stopping'
self._save(state)
raise RunnerContainmentPending(
'runner containment teardown remains pending',
) from exc
if not watchdog_expired.is_set():
watchdog_stop.set()
terminal = self._load_runner_terminal(state)
if terminal is None:
decision, _won = publish_generation_terminal(
paths['terminal'], generation=runner['generation'],
input_sha256=runner['input_sha256'], decision='fenced',
reason='runner_crash',
)
if decision['decision'] == 'completed':
terminal = self._load_runner_terminal(state)
else:
runner['status'] = 'fenced'
self._save(state)
self._drain_runner_events(state)
return self._retry_closed_runner(state, 'runner_crash')
if terminal['decision'] == 'completed':
runner['status'] = 'exited'
self._save(state)
return self._adopt_runner_outcome(state, terminal['outcome'])
runner['status'] = (
'timed_out' if terminal['decision'] == 'timed_out' else 'fenced'
)
self._save(state)
if (
terminal['decision'] == 'timed_out'
and allow_timeout_fallback
and runner['operation'] == 'execute'
):
return self._timeout_bundle(state)
raise RunnerProtocolError('runner generation closed without an adoptable outcome')
finally:
watchdog_stop.set()
if process.poll() is None:
process.kill()
try:
process.wait(timeout=RUNNER_STOP_TIMEOUT_SECONDS)
except subprocess.TimeoutExpired:
pass
def _execute(self, state):
assignment = dict(state['assignment'])
packaged_capabilities = self.package_runtime.get('capabilities')
try:
validate_protocol2_remote_assignment(
assignment, self.compatibility, packaged_capabilities,
)
except (ScanExecutionError, TypeError, ValueError) as exc:
raise WorkerAssignmentCompatibilityError(
'worker assignment is incompatible with this package'
) from exc
adopted = self._adopt_ready_bundle(state)
if adopted is not None:
return adopted
runner = state.get('runner')
if runner is not None:
source = self._runner_paths(runner)['root']
abandoned = os.path.join(
self.work_root, 'abandoned', runner['root_name'],
)
if not os.path.exists(source) and os.path.isdir(abandoned):
return self._retry_closed_runner(
state, 'janitor_transfer_recovery',
)
try:
terminal = self._load_runner_terminal(state)
except RunnerProtocolError:
terminal = None
closed = self._fence_recovered_runner(state, 'malformed_output')
if closed == 'completed':
runner['status'] = 'fenced'
self._save(state)
closed = self._stop_persisted_runner(runner)
if closed == 'unknown':
raise RunnerContainmentPending(
'malformed runner generation containment remains live',
)
return self._retry_closed_runner(state, 'malformed_output')
if terminal is not None and terminal['decision'] == 'completed':
if self._stop_persisted_runner(runner) != 'dead':
raise RunnerContainmentPending(
'completed runner containment remains live during recovery',
)
return self._adopt_runner_outcome(state, terminal['outcome'])
closed = self._fence_recovered_runner(state, 'controller_recovery')
if closed == 'completed':
terminal = self._load_runner_terminal(state)
return self._adopt_runner_outcome(state, terminal['outcome'])
if closed == 'unknown':
raise RunnerContainmentPending(
'recovered assignment runner could not be proven dead',
)
try:
self._drain_runner_events(state)
except RunnerProtocolError:
return self._retry_closed_runner(state, 'malformed_event_tail')
if terminal is not None and terminal['decision'] == 'timed_out':
return self._timeout_bundle(state)
return self._retry_closed_runner(state, 'controller_recovery')
try:
self._new_runner(state)
except RunnerStageTimeout as exc:
if self._event_phase != WorkerPhase.PREPARING:
self._emit(WorkerPhase.PREPARING, state, {
'reason': 'prelaunch_stage_timeout',
})
return self._start_timeout_bundle(
state,
scan_started_at=exc.scan_started_at,
scan_deadline_at=exc.scan_deadline_at,
final_phase=exc.phase,
)
return self._launch_runner(state)
def _terminal(self, state, reservation_id, failure_code=None, detail=None):
if failure_code is not None:
state['phase'] = 'terminal_pending'
report = self._fitted_terminal_report(failure_code, detail)
self._archive_fitted_terminal_diagnostics(report)
encoded = json.dumps(
report, ensure_ascii=True, sort_keys=True, separators=(',', ':'),
).encode('ascii')
if len(encoded) > TERMINAL_REPORT_MAX_BYTES:
raise WorkerClientError('pending terminal report exceeds its byte bound')
state['terminal'] = report
self._save(state)
terminal = dict(state.get('terminal') or {})
if set(terminal) not in (
{'failure_code', 'detail'},
{'failure_code', 'detail', 'diagnostics'},
):
raise WorkerClientError('pending terminal report is invalid')
if self._event_phase not in {
WorkerPhase.UPLOADING, WorkerPhase.AWAITING_RECEIPT,
}:
self._emit(WorkerPhase.UPLOADING, state, {'reason': 'terminal_report'})
try:
awaiting_emitted = False
def awaiting_receipt():
nonlocal awaiting_emitted
if not awaiting_emitted:
self._emit(
WorkerPhase.AWAITING_RECEIPT, state,
{'reason': 'terminal_report'},
)
awaiting_emitted = True
if isinstance(self.api, WorkerHTTPClient):
receipt = self.api.terminal(
reservation_id, terminal,
response_wait_callback=awaiting_receipt,
)
else:
receipt = self.api.terminal(reservation_id, terminal)
awaiting_receipt()
except WorkerHTTPError as exc:
if self._event_phase in {
WorkerPhase.UPLOADING, WorkerPhase.AWAITING_RECEIPT,
}:
self._emit(WorkerPhase.BACKOFF, state, {'reason': 'terminal_retry'})
return self._reconcile_http_failure(state, exc)
except Exception:
if self._event_phase in {
WorkerPhase.UPLOADING, WorkerPhase.AWAITING_RECEIPT,
}:
self._emit(WorkerPhase.BACKOFF, state, {'reason': 'terminal_retry'})
raise
if not self._resolved(state, receipt):
raise WorkerClientError('worker API terminal receipt is not terminal')
return True
def step(self):
state = self._load()
self._resume_events(state)
if state is None:
if not self.claim_enabled:
return False
self._emit(WorkerPhase.CLAIMING)
state = {'phase': 'claiming', 'request_id': secrets.token_hex(16)}
state = self._save(state)
if state.get('assignment') and self._cleanup_recorded_stale(state):
return True
if state.get('phase') == 'claiming':
claim_result = self.api.claim(
state['request_id'], self.compatibility.as_dict(),
)
if claim_result is None:
self._remove_state()
self._emit(WorkerPhase.IDLE)
return False
if set(claim_result) == {'retry_after_seconds', 'reason'}:
retry_after = int(claim_result['retry_after_seconds'])
if not 1 <= retry_after <= 300:
raise WorkerClientError('worker API claim retry delay is invalid')
self.retry_after_seconds = retry_after
self._remove_state()
progress = {
'retry_after_seconds': retry_after,
'next_claim_at': datetime.fromtimestamp(
time.time() + retry_after, timezone.utc,
).isoformat().replace('+00:00', 'Z'),
}
if claim_result['reason'] is not None:
progress['reason'] = claim_result['reason']
self._emit(WorkerPhase.BACKOFF, progress=progress)
return False
if set(claim_result) == {'claim_resolution'}:
self._remove_state()
self._emit(WorkerPhase.IDLE, progress={'reason': 'claim_reconciled'})
return True
assignment = claim_result
self._clear_assignment_diagnostics()
state = {'phase': 'assigned', 'assignment': assignment}
state = self._save(state)
self._emit(WorkerPhase.ASSIGNED, state)
assignment = dict(state.get('assignment') or {})
reservation = dict(assignment.get('reservation') or {})
reservation_id = int(reservation.get('reservation_id') or 0)
if reservation_id <= 0:
raise WorkerClientError('pending assignment has an invalid reservation identity')
try:
status = self.api.status(reservation_id)
except WorkerHTTPError as exc:
if exc.status_code in (404, 409, 410):
return self._mark_stale(state, exc)
raise
if self._resolved(state, status):
return True
if state.get('phase') == 'awaiting_resolution':
return False
if str(status.get('state') or '') != 'scanning' and state.get('phase') != 'bundle_ready':
raise WorkerClientError('pending assignment is no longer executable')
if state.get('phase') == 'terminal_pending':
return self._terminal(state, reservation_id)
if state.get('phase') == 'assigned':
deadline = status.get('expires_at') or reservation.get('remote_expires_at')
if not self._deadline_active(deadline):
raise WorkerClientError('worker assignment deadline has passed')
try:
state = self._execute(state)
except WorkerAssignmentCompatibilityError:
raise
except RunnerContainmentPending:
raise
except RunnerProtocolError as exc:
self._capture_failure_diagnostic(
state, exc, 'runner_protocol_failed', DiagnosticCategory.PROTOCOL,
)
return self._terminal(
state, reservation_id, 'client_process_failed',
CLIENT_PROCESS_FAILURE_DETAIL,
)
except OSError as exc:
self._capture_failure_diagnostic(
state, exc, 'client_storage_failed', DiagnosticCategory.STORAGE,
operation='assignment_execute',
)
return self._terminal(
state, reservation_id, 'client_storage_failed',
CLIENT_STORAGE_FAILURE_DETAIL,
)
except Exception as exc:
self._capture_failure_diagnostic(
state, exc, 'client_process_failed', DiagnosticCategory.INTERNAL,
)
return self._terminal(
state, reservation_id, 'client_process_failed',
CLIENT_PROCESS_FAILURE_DETAIL,
)
if state.get('phase') == 'bundle_ready':
path = bundle_ready_path(self.bundle_root, reservation['bundle_id'])
if not os.path.isfile(path) or os.path.islink(path):
self._capture_failure_diagnostic(
state, WorkerClientError('pending result bundle is missing'),
'client_storage_failed', DiagnosticCategory.STORAGE,
)
return self._terminal(
state, reservation_id, 'client_storage_failed',
'pending result bundle is missing',
)
try:
if self._event_phase in {
WorkerPhase.ASSIGNED, WorkerPhase.BUNDLING, WorkerPhase.BACKOFF,
}:
self._emit(WorkerPhase.UPLOADING, state)
awaiting_emitted = False
def awaiting_receipt():
nonlocal awaiting_emitted
if not awaiting_emitted:
self._emit(WorkerPhase.AWAITING_RECEIPT, state)
awaiting_emitted = True
if isinstance(self.api, WorkerHTTPClient):
receipt = self.api.upload(
reservation_id, path,
response_wait_callback=awaiting_receipt,
)
else:
receipt = self.api.upload(reservation_id, path)
awaiting_receipt()
except WorkerHTTPError as exc:
if self._event_phase in {WorkerPhase.UPLOADING, WorkerPhase.AWAITING_RECEIPT}:
self._emit(WorkerPhase.BACKOFF, state, {'reason': 'upload_retry'})
return self._reconcile_http_failure(state, exc)
except Exception:
if self._event_phase in {WorkerPhase.UPLOADING, WorkerPhase.AWAITING_RECEIPT}:
self._emit(WorkerPhase.BACKOFF, state, {'reason': 'upload_retry'})
raise
if not self._resolved(state, receipt):
raise WorkerClientError('worker API upload receipt is not terminal')
return True
raise WorkerClientError('pending slot phase is invalid')
def persisted_slot_ids(state_dir):
slot_ids = set()
entries = 0
with os.scandir(state_dir) as iterator:
for entry in iterator:
entries += 1
if entries > 4096:
raise WorkerClientError('worker state directory exceeds its entry bound')
match = re.fullmatch(r'slot-(0|[1-9][0-9]*)\.json', entry.name)
if not match:
continue
if entry.is_symlink() or not entry.is_file(follow_symlinks=False):
raise WorkerClientError('worker slot state is not a regular file')
slot_id = int(match.group(1))
if slot_id > 100000:
raise WorkerClientError('worker slot identity exceeds its bound')
slot_ids.add(slot_id)
return slot_ids
def persisted_runner_root_names(state_dir):
roots = set()
for slot_id in persisted_slot_ids(state_dir):
path = os.path.join(state_dir, f'slot-{slot_id}.json')
try:
state = WorkerSlot._slot_state(
read_private_json(path, max_bytes=MAX_PENDING_BYTES),
slot_id=slot_id,
)
except FileNotFoundError:
continue
runner = state.get('runner')
if runner is not None:
roots.add(runner['root_name'])
roots.add(f"abandoned/{runner['root_name']}")
return roots
def default_worker_paths(module_path=None, *, platform_name=None, environ=None, home=None):
platform_name = str(platform_name or os.name)
path_module = ntpath if platform_name == 'nt' else posixpath
environment = os.environ if environ is None else environ
package_root = path_module.dirname(path_module.dirname(path_module.abspath(
module_path or __file__,
)))
home = path_module.abspath(home or os.path.expanduser('~'))
if platform_name == 'nt':
state_base = str(environment.get('LOCALAPPDATA') or '')
if state_base and not path_module.isabs(state_base):
raise ValueError('LOCALAPPDATA must be absolute')
if not state_base:
state_base = path_module.join(home, 'AppData', 'Local')
state_dir = path_module.join(state_base, 'TRUF', 'RemoteWorker')
bundle_dir = path_module.join(state_dir, 'bundles')
work_dir = path_module.join(state_dir, 'work')
else:
state_base = str(environment.get('XDG_STATE_HOME') or '')
data_base = str(environment.get('XDG_DATA_HOME') or '')
if state_base and not path_module.isabs(state_base):
raise ValueError('XDG_STATE_HOME must be absolute')
if data_base and not path_module.isabs(data_base):
raise ValueError('XDG_DATA_HOME must be absolute')
state_base = state_base or path_module.join(home, '.local', 'state')
data_base = data_base or path_module.join(home, '.local', 'share')
state_dir = path_module.join(state_base, 'truf', 'remote-worker')
bundle_dir = path_module.join(data_base, 'truf', 'remote-worker', 'bundles')
work_dir = path_module.join(data_base, 'truf', 'remote-worker', 'work')
return {
'package_manifest': path_module.join(package_root, 'worker-package.json'),
'state_dir': state_dir,
'bundle_dir': bundle_dir,
'work_dir': work_dir,
}
def run_client(
args, *, drain_event=None, stop_event=None, event_callback=None,
terminal_callback=None, log_callback=None, acquire_lock=True,
started_callback=None, diagnostic_callback=None,
progress_event_reader=None,
):
package = verify_worker_package(args.package_manifest)
manifest = package.pop('manifest')
compatibility = WorkerBuildCompatibility.from_mapping(
package.pop('build_compatibility'),
)
package.pop('runtime_trees')
package['capabilities'] = tuple(
(
capability['source'], capability['platform'],
capability['planning_kind'],
)
for capability in manifest['capabilities']
)
package['bootstrap_path'] = os.path.join(
os.path.dirname(os.path.abspath(args.package_manifest)),
'app', 'remote_worker_bootstrap.py',
)
state_dir = ensure_private_directory(os.path.abspath(args.state_dir), reject_reparse=True)
bundle_root = ensure_private_directory(os.path.abspath(args.bundle_dir), reject_reparse=True)
work_root = ensure_private_directory(os.path.abspath(args.work_dir), reject_reparse=True)
for name in ('tmp', 'ready', 'quarantine'):
ensure_private_directory(os.path.join(bundle_root, name), reject_reparse=True)
scanner.scan_config.work_dir = work_root
api = WorkerHTTPClient(args.server, args.token, args.http_timeout)
stopping = stop_event or threading.Event()
draining = drain_event or threading.Event()
progress_outbox = (
ProgressOutbox(api, state_dir, progress_event_reader)
if progress_event_reader is not None else None
)
progress_stopping = threading.Event()
progress_thread = (
threading.Thread(
target=progress_outbox.run,
args=(progress_stopping,),
name='worker-progress-outbox', daemon=True,
)
if progress_outbox is not None else None
)
configured_slot_ids = set(range(args.parallelism))
slot_ids = sorted(configured_slot_ids | persisted_slot_ids(state_dir))
start_gate = threading.Event()
ready_events = {slot_id: threading.Event() for slot_id in slot_ids}
startup_errors = {}
startup_lock = threading.Lock()
def loop(slot_id):
try:
slot = WorkerSlot(
slot_id, api, compatibility.as_dict(), state_dir, bundle_root,
work_root=work_root,
package_runtime=package,
claim_enabled=slot_id in configured_slot_ids,
event_callback=event_callback,
terminal_callback=terminal_callback,
diagnostic_callback=diagnostic_callback,
)
except BaseException as exc:
with startup_lock:
startup_errors[slot_id] = exc
ready_events[slot_id].set()
return
ready_events[slot_id].set()
start_gate.wait()
while not stopping.is_set():
if draining.is_set():
slot.claim_enabled = False
if not slot.claim_enabled and not os.path.exists(slot.state_path):
if hasattr(slot, 'retire'):
slot.retire(
'graceful_drain' if draining.is_set()
else 'lowered_parallelism_recovery_complete'
)
return
try:
worked = slot.step()
delay = 0 if worked else (
slot.retry_after_seconds
if slot.retry_after_seconds is not None else args.poll_seconds
)
slot.retry_after_seconds = None
except RunnerContainmentPending:
message = f'worker slot {slot_id}: containment authority remains pending'
print(message, flush=True)
if log_callback is not None:
log_callback(message)
return
except Exception as exc:
print(
f'worker slot {slot_id}: {safe_worker_error_summary(exc)}',
flush=True,
)
if log_callback is not None:
log_callback(
f'worker slot {slot_id}: {safe_worker_error_summary(exc)}'
)
delay = args.error_delay_seconds
stopping.wait(max(0.1, float(delay)))
threads = [
threading.Thread(target=loop, args=(slot_id,), name=f'worker-slot-{slot_id}', daemon=True)
for slot_id in slot_ids
]
active_threads = threads[:args.parallelism]
lock = (
PrivateFileLock(os.path.join(state_dir, 'remote-worker.lock'))
if acquire_lock else nullcontext()
)
with lock:
forced_stop = False
unexpected_exit = False
try:
if progress_thread is not None:
progress_thread.start()
for thread in threads:
thread.start()
startup_deadline = time.monotonic() + 5.0
for slot_id in slot_ids:
remaining = startup_deadline - time.monotonic()
if remaining <= 0 or not ready_events[slot_id].wait(remaining):
raise WorkerClientError('worker slot startup timed out')
if startup_errors or not all(thread.is_alive() for thread in threads):
raise WorkerClientError('worker slot startup failed')
if started_callback is not None:
started_callback()
start_gate.set()
while True:
if stopping.is_set():
forced_stop = True
break
monitored = threads if (draining.is_set() or stopping.is_set()) else active_threads
if not monitored or not all(thread.is_alive() for thread in monitored):
if not (draining.is_set() or stopping.is_set()) or not any(
thread.is_alive() for thread in threads
):
unexpected_exit = not (draining.is_set() or stopping.is_set())
break
time.sleep(0.5)
except KeyboardInterrupt:
forced_stop = True
stopping.set()
finally:
stopping.set()
start_gate.set()
if not forced_stop:
for thread in threads:
thread.join(timeout=5)
if progress_thread is not None:
progress_shutdown_started = time.monotonic()
progress_stopping.set()
api.cancel_progress_requests()
progress_thread.join(timeout=PROGRESS_REQUEST_TIMEOUT_SECONDS + 0.5)
if progress_thread.is_alive():
raise WorkerClientError(
'progress publisher exceeded its absolute shutdown bound'
)
if not forced_stop:
progress_outbox.drain(max(
0.0,
PROGRESS_FINAL_DRAIN_SECONDS
- (time.monotonic() - progress_shutdown_started),
))
return 2 if forced_stop else (1 if unexpected_exit else 0)
def parse_args(argv=None):
defaults = default_worker_paths()
parser = argparse.ArgumentParser(description='Trusted TRUF remote scan worker')
parser.add_argument('--server', required=True)
parser.add_argument('--token', required=True)
parser.add_argument('--parallelism', type=int, default=1)
parser.set_defaults(
package_manifest=defaults['package_manifest'],
state_dir=defaults['state_dir'],
bundle_dir=defaults['bundle_dir'],
work_dir=defaults['work_dir'],
poll_seconds=5.0,
error_delay_seconds=15.0,
http_timeout=120,
)
args = parser.parse_args(argv)
if not 1 <= args.parallelism <= 128:
parser.error('--parallelism must be between 1 and 128')
return args
def main(argv=None):
run_client(parse_args(argv))
if __name__ == '__main__':
main()