import hmac import json import os import re import secrets import signal import socket import socketserver import struct import subprocess import sys import threading import time from datetime import datetime, timezone from process_identity import current_process_identity, exact_process_identity_state, open_process from remote_worker_client import persisted_runner_root_names, persisted_slot_ids, run_client from runtime_security import ( PrivateFileLock, atomic_write_private_json, canonical_path, durable_unlink, ensure_private_directory, read_private_json, write_private_json_exclusive, ) from worker_contracts import WORKER_EVENT_SCHEMA, WorkerPhase from worker_local_state import ( WorkerLocalState, prepare_progress_outbox_cursor, utc_now, ) from worker_package import verify_worker_package, worker_package_manifest_sha256 from worker_assignment_runner import cleanup_abandoned_runner_roots INSTANCE_SCHEMA = 1 STARTUP_SCHEMA = 1 SHUTDOWN_SCHEMA = 1 CONTROL_SCHEMA = 1 CONTROL_MAX_REQUEST_BYTES = 64 * 1024 CONTROL_MAX_RESPONSE_BYTES = 2 * 1024 * 1024 CONTROL_TIMEOUT_SECONDS = 5.0 CONTROL_MAX_WORKERS = 8 PROJECTION_SCHEMA = 1 class WorkerSupervisorError(RuntimeError): pass class WorkerAlreadyRunning(WorkerSupervisorError): pass class WorkerInstanceUnverifiable(WorkerSupervisorError): pass def instance_path(state_dir): return os.path.join(os.path.abspath(state_dir), 'control', 'worker.instance.json') def shutdown_path(state_dir): return os.path.join(os.path.abspath(state_dir), 'control', 'worker.exit.json') def worker_lock_path(state_dir): return os.path.join(os.path.abspath(state_dir), 'remote-worker.lock') 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 WorkerSupervisorError('control JSON is invalid') from exc def _decode_canonical(payload, maximum): if type(payload) is not bytes or not payload or len(payload) > maximum: raise WorkerSupervisorError('control JSON is invalid or oversized') def reject_duplicate(pairs): value = {} for key, item in pairs: if key in value: raise WorkerSupervisorError('control JSON contains duplicate fields') value[key] = item return value try: value = json.loads( payload.decode('ascii'), object_pairs_hook=reject_duplicate, parse_constant=lambda _value: (_ for _ in ()).throw( WorkerSupervisorError('control JSON constant is invalid') ), ) except WorkerSupervisorError: raise except (UnicodeDecodeError, json.JSONDecodeError, TypeError, ValueError) as exc: raise WorkerSupervisorError('control JSON is invalid') from exc if not isinstance(value, dict) or not hmac.compare_digest(_canonical_json(value), payload): raise WorkerSupervisorError('control JSON is not a canonical object') return value def encode_frame(value, maximum=CONTROL_MAX_RESPONSE_BYTES): payload = _canonical_json(value) if not payload or len(payload) > int(maximum): raise WorkerSupervisorError('control frame exceeds its bound') return struct.pack('!I', len(payload)) + payload def _receive_exact(sock, size): chunks = [] remaining = int(size) while remaining: chunk = sock.recv(remaining) if not chunk: raise WorkerSupervisorError('control frame ended early') chunks.append(chunk) remaining -= len(chunk) return b''.join(chunks) def receive_frame(sock, maximum=CONTROL_MAX_REQUEST_BYTES): header = _receive_exact(sock, 4) length = struct.unpack('!I', header)[0] if length < 2 or length > int(maximum): raise WorkerSupervisorError('control frame length is invalid') return _decode_canonical(_receive_exact(sock, length), int(maximum)) def _package_identity(package_manifest): verified = verify_worker_package(package_manifest) manifest = verified['manifest'] return { 'schema': manifest['schema'], 'manifest_sha256': worker_package_manifest_sha256(manifest), 'code_manifest_sha256': verified['code_manifest_sha256'], 'platform_tag': manifest['platform_tag'], }, { 'worker_protocol': manifest['protocol_version'], 'bundle_format': manifest['bundle_format_version'], 'event': WORKER_EVENT_SCHEMA, 'control': CONTROL_SCHEMA, 'projection': PROJECTION_SCHEMA, } def _runtime_identity(foreground): return { 'python': '.'.join(str(item) for item in sys.version_info[:3]), 'platform': sys.platform, 'executable': canonical_path(sys.executable), 'mode': 'foreground' if foreground else 'detached', } def build_instance_record( *, instance_id, token, identity, package, protocol, control_port, foreground, started_at=None, lifecycle='running', ): return validate_instance_record({ 'schema': INSTANCE_SCHEMA, 'instance_id': str(instance_id), 'token': str(token), 'pid': int(identity.pid), 'process_creation_time': str(identity.creation_time), 'executable': canonical_path(identity.executable), 'package': dict(package), 'runtime': _runtime_identity(foreground), 'protocol': dict(protocol), 'control': {'host': '127.0.0.1', 'port': int(control_port)}, 'started_at': started_at or utc_now(), 'lifecycle': str(lifecycle), }) def validate_instance_record(value): fields = { 'schema', 'instance_id', 'token', 'pid', 'process_creation_time', 'executable', 'package', 'runtime', 'protocol', 'control', 'started_at', 'lifecycle', } if not isinstance(value, dict) or set(value) != fields or value.get('schema') != INSTANCE_SCHEMA: raise WorkerSupervisorError('worker instance record shape is invalid') if not isinstance(value.get('instance_id'), str) or not value['instance_id']: raise WorkerSupervisorError('worker instance identity is invalid') if not isinstance(value.get('token'), str) or not 32 <= len(value['token']) <= 512: raise WorkerSupervisorError('worker instance token is invalid') if type(value.get('pid')) is not int or value['pid'] <= 0: raise WorkerSupervisorError('worker instance PID is invalid') if not isinstance(value.get('process_creation_time'), str) or not value['process_creation_time']: raise WorkerSupervisorError('worker process creation identity is invalid') if not isinstance(value.get('executable'), str) or not value['executable']: raise WorkerSupervisorError('worker executable identity is invalid') package = value.get('package') if not isinstance(package, dict) or set(package) != { 'schema', 'manifest_sha256', 'code_manifest_sha256', 'platform_tag', }: raise WorkerSupervisorError('worker package identity is invalid') if any(not isinstance(package[name], str) or not package[name] for name in ( 'manifest_sha256', 'code_manifest_sha256', 'platform_tag', )) or type(package['schema']) is not int: raise WorkerSupervisorError('worker package identity fields are invalid') if ( re.fullmatch(r'[0-9a-f]{64}', package['manifest_sha256']) is None or re.fullmatch(r'[0-9a-f]{64}', package['code_manifest_sha256']) is None or re.fullmatch(r'(windows|linux)-(x86_64|aarch64)', package['platform_tag']) is None ): raise WorkerSupervisorError('worker package identity fields are invalid') runtime = value.get('runtime') if not isinstance(runtime, dict) or set(runtime) != { 'python', 'platform', 'executable', 'mode', } or runtime.get('mode') not in {'foreground', 'detached'}: raise WorkerSupervisorError('worker runtime identity is invalid') if any(not isinstance(runtime.get(name), str) or not runtime[name] for name in ( 'python', 'platform', 'executable', )): raise WorkerSupervisorError('worker runtime identity fields are invalid') protocol = value.get('protocol') if not isinstance(protocol, dict) or set(protocol) != { 'worker_protocol', 'bundle_format', 'event', 'control', 'projection', } or any(type(item) is not int for item in protocol.values()): raise WorkerSupervisorError('worker protocol identity is invalid') if any(item <= 0 for item in protocol.values()): raise WorkerSupervisorError('worker protocol identity is invalid') control = value.get('control') if not isinstance(control, dict) or set(control) != {'host', 'port'}: raise WorkerSupervisorError('worker control identity is invalid') if control.get('host') != '127.0.0.1' or type(control.get('port')) is not int or not 0 < control['port'] <= 65535: raise WorkerSupervisorError('worker control endpoint is invalid') if value.get('lifecycle') not in {'starting', 'running', 'draining'}: raise WorkerSupervisorError('worker lifecycle state is invalid') if not isinstance(value.get('started_at'), str) or not value['started_at'].endswith('Z'): raise WorkerSupervisorError('worker startup timestamp is invalid') normalized = dict(value) normalized['executable'] = canonical_path(value['executable']) normalized['runtime'] = dict(runtime) normalized['runtime']['executable'] = canonical_path(runtime['executable']) if normalized['runtime']['executable'] != normalized['executable']: raise WorkerSupervisorError('worker runtime executable identity is inconsistent') normalized['package'] = dict(package) normalized['protocol'] = dict(protocol) normalized['control'] = dict(control) return normalized def load_instance(state_dir): return validate_instance_record(read_private_json(instance_path(state_dir))) def public_instance(record): record = validate_instance_record(record) return { 'schema': record['schema'], 'instance_id': record['instance_id'], 'pid': record['pid'], 'process_creation_time': record['process_creation_time'], 'executable': record['executable'], 'package': record['package'], 'runtime': record['runtime'], 'protocol': record['protocol'], 'control': record['control'], 'started_at': record['started_at'], 'lifecycle': record['lifecycle'], } def _remove_exact_stale(state_dir, record, identity_state=exact_process_identity_state): if identity_state( record['pid'], record['process_creation_time'], record['executable'], ) not in {'dead', 'reused'}: return False lock = PrivateFileLock(worker_lock_path(state_dir)) try: lock.acquire() except OSError: return False try: current = load_instance(state_dir) if not hmac.compare_digest(current['instance_id'], record['instance_id']): return False if identity_state( current['pid'], current['process_creation_time'], current['executable'], ) not in {'dead', 'reused'}: return False durable_unlink(instance_path(state_dir)) return True except (OSError, ValueError, WorkerSupervisorError): return False finally: lock.release() def classify_instance( state_dir, *, identity_state=exact_process_identity_state, request=None, remove_stale=False, ): path = instance_path(state_dir) if not os.path.exists(path): return {'state': 'stopped', 'instance': None, 'detail': 'no instance record'} try: record = load_instance(state_dir) except (OSError, ValueError, WorkerSupervisorError) as exc: return {'state': 'unverifiable', 'instance': None, 'detail': str(exc)} state = identity_state( record['pid'], record['process_creation_time'], record['executable'], ) if state in {'dead', 'reused'}: removed = ( _remove_exact_stale(state_dir, record, identity_state=identity_state) if remove_stale else False ) return { 'state': 'stale', 'instance': public_instance(record), 'detail': f'process identity is {state}', 'reason': f'process_{state}', 'removable': True, 'removed': removed, } if state != 'alive': return { 'state': 'unverifiable', 'instance': public_instance(record), 'detail': 'process identity could not be verified', } try: response = (request or send_control_request)(record, 'handshake', {}) except (OSError, TimeoutError, ValueError, WorkerSupervisorError) as exc: return { 'state': 'stale', 'instance': public_instance(record), 'detail': f'live process control handshake failed: {type(exc).__name__}', 'reason': 'control_handshake_failed', 'removable': False, 'removed': False, } if not isinstance(response, dict) or set(response) != { 'schema', 'instance_id', 'lifecycle', 'sequence', } or ( response.get('schema') != 1 or response.get('instance_id') != record['instance_id'] or response.get('lifecycle') not in {'starting', 'running', 'draining'} or type(response.get('sequence')) is not int or response['sequence'] < 0 ): return { 'state': 'stale', 'instance': public_instance(record), 'detail': 'control handshake response is invalid', 'reason': 'control_handshake_invalid', 'removable': False, 'removed': False, } lifecycle = response['lifecycle'] public = public_instance(record) public['lifecycle'] = lifecycle return { 'state': lifecycle, 'instance': public, 'detail': f'verified process and {lifecycle} control handshake', 'record': record, 'handshake': response, } def capture_spawned_process_identity(process, *, opener=open_process): if process is None or type(getattr(process, 'pid', None)) is not int or process.pid <= 0: raise WorkerSupervisorError('spawned worker process handle is invalid') retained = opener(process.pid) try: identity = retained.identity if process.poll() is not None: raise WorkerSupervisorError('spawned worker exited during identity capture') return { 'pid': identity.pid, 'creation_time': identity.creation_time, 'executable': identity.executable, } finally: retained.close() def terminate_spawned_process( process, identity, *, identity_state=exact_process_identity_state, timeout=5.0, ): if process is None or process.poll() is not None: return True if not isinstance(identity, dict) or set(identity) != { 'pid', 'creation_time', 'executable', }: raise WorkerSupervisorError('spawned worker exact identity is unavailable') state = identity_state( identity['pid'], identity['creation_time'], identity['executable'], ) if state != 'alive' or int(process.pid) != int(identity['pid']): if process.poll() is not None: return True raise WorkerSupervisorError('spawned worker exact identity is no longer retained') process.terminate() try: process.wait(timeout=max(0.1, float(timeout))) except subprocess.TimeoutExpired: state = identity_state( identity['pid'], identity['creation_time'], identity['executable'], ) if state != 'alive': return process.poll() is not None process.kill() process.wait(timeout=max(0.1, float(timeout))) return process.poll() is not None def send_control_request(record, action, parameters=None, timeout=CONTROL_TIMEOUT_SECONDS): record = validate_instance_record(record) request = { 'schema': CONTROL_SCHEMA, 'instance_id': record['instance_id'], 'token': record['token'], 'action': str(action), 'parameters': dict(parameters or {}), } deadline = time.monotonic() + max(0.1, float(timeout)) remaining = lambda: max(0.01, deadline - time.monotonic()) with socket.create_connection( (record['control']['host'], record['control']['port']), timeout=remaining(), ) as connection: connection.settimeout(remaining()) connection.sendall(encode_frame(request, CONTROL_MAX_REQUEST_BYTES)) response = receive_frame(connection, CONTROL_MAX_RESPONSE_BYTES) expected = ( {'schema', 'instance_id', 'ok', 'result'} if response.get('ok') is True else {'schema', 'instance_id', 'ok', 'error'} ) if ( set(response) != expected or response.get('schema') != CONTROL_SCHEMA or response.get('instance_id') != record['instance_id'] or type(response.get('ok')) is not bool ): raise WorkerSupervisorError('control response shape is invalid') if not response['ok']: if not isinstance(response.get('error'), str): raise WorkerSupervisorError('control response error is invalid') raise WorkerSupervisorError(response['error']) if not isinstance(response.get('result'), dict): raise WorkerSupervisorError('control response result is invalid') return response['result'] class _ControlHandler(socketserver.BaseRequestHandler): def handle(self): self.request.settimeout(CONTROL_TIMEOUT_SECONDS) try: request = receive_frame(self.request, CONTROL_MAX_REQUEST_BYTES) response = self.server.runtime.control_request(request) except Exception: response = self.server.runtime.control_error('invalid control request') try: self.request.sendall(encode_frame(response, CONTROL_MAX_RESPONSE_BYTES)) except OSError: pass class _ControlServer(socketserver.ThreadingMixIn, socketserver.TCPServer): allow_reuse_address = False daemon_threads = True request_queue_size = 16 def __init__(self, address, runtime): self.runtime = runtime self._workers = threading.BoundedSemaphore(CONTROL_MAX_WORKERS) super().__init__(address, _ControlHandler) def server_bind(self): if os.name == 'nt' and hasattr(socket, 'SO_EXCLUSIVEADDRUSE'): self.socket.setsockopt(socket.SOL_SOCKET, socket.SO_EXCLUSIVEADDRUSE, 1) super().server_bind() def process_request(self, request, client_address): if not self._workers.acquire(blocking=False): try: request.sendall(encode_frame( self.runtime.control_error('control worker limit reached'), CONTROL_MAX_RESPONSE_BYTES, )) finally: self.shutdown_request(request) return try: super().process_request(request, client_address) except BaseException: self._workers.release() raise def process_request_thread(self, request, client_address): try: super().process_request_thread(request, client_address) finally: self._workers.release() class WorkerSupervisor: def __init__( self, args, *, foreground=True, local_state_factory=WorkerLocalState, monotonic=time.monotonic, wall_time=time.time, retention_interval_seconds=300.0, hard_exit_hook=os._exit, ): self.args = args self.foreground = bool(foreground) self.state_dir = ensure_private_directory(os.path.abspath(args.state_dir), reject_reparse=True) self.control_dir = os.path.join(self.state_dir, 'control') self.local = None self._local_state_factory = local_state_factory self._monotonic = monotonic self._wall_time = wall_time self._retention_interval = max(1.0, float(retention_interval_seconds)) self._hard_exit_hook = hard_exit_hook self.instance_id = secrets.token_urlsafe(24) self.token = secrets.token_urlsafe(48) self.identity = current_process_identity() self.package, self.protocol = _package_identity(args.package_manifest) self.started_at = utc_now() self.drain_event = threading.Event() self.stop_event = threading.Event() self.drain_deadline = None self._drain_deadline_monotonic = None self._deadline_escalated = False self.record = None self.server = None self.server_thread = None self._lock = PrivateFileLock(worker_lock_path(self.state_dir)) self._record_lock = threading.RLock() self._startup_ready = False self._maintenance_stop = threading.Event() self._maintenance_wakeup = threading.Event() self._maintenance_thread = None self._maintenance_started = False self._next_retention_cleanup = None self._server_started = False self._hard_exit_invoked = False def _initialize_local_state(self): if self.local is None: self.local = self._local_state_factory( self.state_dir, log_bytes=getattr(self.args, 'log_bytes', 2 * 1024 * 1024), log_files=getattr(self.args, 'log_files', 5), retention_days=getattr(self.args, 'retention_days', 30), retention_bytes=getattr(self.args, 'retention_bytes', 1024 * 1024 * 1024), ) return self.local def _public_worker(self): projection = self.local.snapshot() configured = set(range(int(self.args.parallelism))) try: slots = configured | persisted_slot_ids(self.state_dir) except OSError: slots = configured return { 'state': ( 'draining' if self.drain_event.is_set() else (self.record or {}).get('lifecycle', 'starting') ), 'parallelism': int(self.args.parallelism), 'slot_cap': len(configured), 'configured_slots': len(configured), 'recovery_slots': len(slots - configured), 'started_at': self.started_at, 'drain_deadline_at': self.drain_deadline, 'aggregate': projection['aggregate'], } def snapshot(self): return { 'schema': PROJECTION_SCHEMA, 'instance': public_instance(self.record), 'worker': self._public_worker(), 'projection': self.local.snapshot(), 'retention': self.local.retention_usage({ 'bundles': self.args.bundle_dir, 'work': self.args.work_dir, }), } def control_error(self, message): return { 'schema': CONTROL_SCHEMA, 'instance_id': self.instance_id, 'ok': False, 'error': str(message), } def control_ok(self, result): return { 'schema': CONTROL_SCHEMA, 'instance_id': self.instance_id, 'ok': True, 'result': dict(result), } def control_request(self, request): if not isinstance(request, dict) or set(request) != { 'schema', 'instance_id', 'token', 'action', 'parameters', }: return self.control_error('control request shape is invalid') if ( request.get('schema') != CONTROL_SCHEMA or not isinstance(request.get('parameters'), dict) or not hmac.compare_digest(str(request.get('instance_id') or ''), self.instance_id) or not hmac.compare_digest(str(request.get('token') or ''), self.token) ): return self.control_error('control authentication failed') action = request.get('action') parameters = request['parameters'] if action == 'handshake' and not parameters: return self.control_ok({ 'schema': 1, 'instance_id': self.instance_id, 'lifecycle': ( 'draining' if self.drain_event.is_set() else (self.record or {}).get('lifecycle', 'starting') ), 'sequence': self.local.snapshot()['sequence'], }) if action == 'snapshot' and not parameters: return self.control_ok(self.snapshot()) if action == 'events' and set(parameters) == {'after_sequence', 'limit'}: try: events = self.local.events_after( parameters['after_sequence'], parameters['limit'], ) except (TypeError, ValueError): return self.control_error('event request bounds are invalid') return self.control_ok({ 'schema': 1, 'events': events, 'last_sequence': self.local.snapshot()['sequence'], }) if action == 'stop' and set(parameters) == {'timeout_seconds'}: try: timeout = float(parameters['timeout_seconds']) except (TypeError, ValueError, OverflowError): return self.control_error('stop timeout is invalid') if not 0.1 <= timeout <= 3600: return self.control_error('stop timeout is invalid') self.request_drain(timeout) return self.control_ok({ 'schema': 1, 'accepted': True, 'drain_deadline_at': self.drain_deadline, 'slots': self.local.snapshot()['slots'], }) return self.control_error('control action is invalid') def _update_lifecycle(self, lifecycle): with self._record_lock: if self.record is None or self.record['lifecycle'] == lifecycle: return self.record = dict(self.record) self.record['lifecycle'] = lifecycle self.record = validate_instance_record(self.record) atomic_write_private_json(instance_path(self.state_dir), self.record) def request_drain(self, timeout=30.0): timeout = max(0.1, float(timeout)) monotonic_deadline = self._monotonic() + timeout deadline = datetime.fromtimestamp( self._wall_time() + timeout, timezone.utc, ).isoformat(timespec='milliseconds').replace('+00:00', 'Z') if ( self._drain_deadline_monotonic is None or monotonic_deadline < self._drain_deadline_monotonic ): self._drain_deadline_monotonic = monotonic_deadline self.drain_deadline = deadline self.drain_event.set() self._update_lifecycle('draining') self.local.log('graceful drain requested') self._maintenance_wakeup.set() def request_interrupt(self, signum): if self.drain_event.is_set(): self._deadline_escalated = True self.stop_event.set() self.local.log(f'interrupt {signum} forced supervisor exit') self._maintenance_wakeup.set() return 'forced' self.request_drain(30.0) self.local.log(f'interrupt {signum} requested graceful drain') return 'draining' def _maintenance_tick(self): now = self._monotonic() if ( self._drain_deadline_monotonic is not None and now >= self._drain_deadline_monotonic and not self.stop_event.is_set() ): self._deadline_escalated = True self.stop_event.set() self.local.log('graceful drain deadline expired; stopping controller') if self._next_retention_cleanup is None or now >= self._next_retention_cleanup: self.local.cleanup_retention({ 'bundles': self.args.bundle_dir, 'work': self.args.work_dir, }) if os.path.isdir(self.args.work_dir): cleanup_abandoned_runner_roots( self.args.work_dir, minimum_age_sec=60, active_root_names=persisted_runner_root_names(self.state_dir), ) self._next_retention_cleanup = now + self._retention_interval def _maintenance_loop(self): while not self._maintenance_stop.is_set(): try: self._maintenance_tick() except Exception as exc: self.local.log(f'worker retention maintenance failed: {type(exc).__name__}') wait_seconds = 1.0 if self._drain_deadline_monotonic is not None: wait_seconds = min( wait_seconds, max(0.01, self._drain_deadline_monotonic - self._monotonic()), ) self._maintenance_wakeup.wait(wait_seconds) self._maintenance_wakeup.clear() def _start_maintenance(self): self._next_retention_cleanup = self._monotonic() + self._retention_interval self._maintenance_thread = threading.Thread( target=self._maintenance_loop, name='worker-maintenance', daemon=True, ) self._maintenance_thread.start() self._maintenance_started = True def _stop_maintenance(self): self._maintenance_stop.set() self._maintenance_wakeup.set() if self._maintenance_started and self._maintenance_thread is not None: self._maintenance_thread.join(timeout=5) self._maintenance_started = False def _event(self, value): return self.local.emit_phase( self.instance_id, value['slot_id'], value['phase'], reservation_id=value.get('reservation_id'), source=value.get('source'), scan_deadline_at=value.get('scan_deadline_at'), assignment_deadline_at=value.get('assignment_deadline_at'), progress=value.get('progress'), timestamp=value.get('timestamp'), phase_started_at=value.get('phase_started_at'), ) def _diagnostic(self, envelope, full_materials=None): return self.local.archive_diagnostic(envelope, full_materials) def _terminal(self, value): completed = value.get('completed_at') or utc_now() timeline = self.local.assignment_timeline( value['slot_id'], value['reservation_id'], ) phase_durations = {} for index, event in enumerate(timeline): end = timeline[index + 1]['timestamp'] if index + 1 < len(timeline) else completed try: started_value = datetime.fromisoformat(event['timestamp'].replace('Z', '+00:00')) ended_value = datetime.fromisoformat(end.replace('Z', '+00:00')) duration = max(0.0, (ended_value - started_value).total_seconds()) except ValueError: duration = 0.0 phase_durations[event['phase']] = round( phase_durations.get(event['phase'], 0.0) + duration, 6, ) diagnostics = { reference['diagnostic_uid']: reference for reference in self.local.diagnostic_references(value['reservation_id']) } diagnostics.update({ reference['diagnostic_uid']: reference for reference in value.get('diagnostics') or [] }) record = { 'history_id': str(value['history_id']), 'instance_id': self.instance_id, 'slot_id': int(value['slot_id']), 'reservation_id': int(value['reservation_id']), 'source': value.get('source'), 'outcome': str(value['outcome']), 'receipt': dict(value.get('receipt') or {}), 'started_at': value.get('started_at'), 'completed_at': completed, 'duration_seconds': value.get('duration_seconds'), 'first_sequence': timeline[0]['sequence'] if timeline else value.get('first_sequence'), 'last_sequence': timeline[-1]['sequence'] if timeline else self.local.snapshot()['sequence'], 'diagnostics': [diagnostics[key] for key in sorted(diagnostics)], 'timeline': timeline, 'phase_durations': phase_durations, } self.local.append_history(record) def _publish_instance(self): self._initialize_local_state() path = instance_path(self.state_dir) if os.path.exists(path): previous = load_instance(self.state_dir) state = exact_process_identity_state( previous['pid'], previous['process_creation_time'], previous['executable'], ) if state not in {'dead', 'reused'}: raise WorkerInstanceUnverifiable('existing worker instance is not exactly stale') durable_unlink(path) if os.path.exists(shutdown_path(self.state_dir)): durable_unlink(shutdown_path(self.state_dir)) try: self.server = _ControlServer(('127.0.0.1', 0), self) self.record = build_instance_record( instance_id=self.instance_id, token=self.token, identity=self.identity, package=self.package, protocol=self.protocol, control_port=self.server.server_address[1], foreground=self.foreground, started_at=self.started_at, lifecycle='starting', ) write_private_json_exclusive(path, self.record) self.server_thread = threading.Thread( target=self.server.serve_forever, name='worker-control', daemon=True, ) self.server_thread.start() self._server_started = True except BaseException: self._remove_instance() if self.server is not None: self.server.server_close() self.server = None self.server_thread = None self.record = None raise def _remove_instance(self): path = instance_path(self.state_dir) try: current = load_instance(self.state_dir) if hmac.compare_digest(current['instance_id'], self.instance_id): durable_unlink(path) except (OSError, ValueError, WorkerSupervisorError): pass def _write_receipt(self, exit_code, *, drained=None): if drained is None: try: pending_slots = persisted_slot_ids(self.state_dir) except OSError: pending_slots = {-1} drained = not pending_slots atomic_write_private_json(shutdown_path(self.state_dir), { 'schema': SHUTDOWN_SCHEMA, 'instance_id': self.instance_id, 'completed_at': utc_now(), 'exit_code': int(exit_code), 'drained': bool(drained), 'last_sequence': self.local.snapshot()['sequence'], }) def _shutdown_control(self): if self.server is not None: if self._server_started: self.server.shutdown() self.server.server_close() if self._server_started and self.server_thread is not None: self.server_thread.join(timeout=5) self._server_started = False def run(self, startup_file=None, launch_nonce=None): try: self._lock.acquire() except OSError as exc: if startup_file: write_startup_result(startup_file, launch_nonce, 'already_running', error='singleton lock is held') raise WorkerAlreadyRunning('another worker supervisor is already active') from exc previous_handlers = {} exit_code = 1 try: self._initialize_local_state() self._publish_instance() prepare_progress_outbox_cursor(self.state_dir, create=True) self._start_maintenance() self.local.log('worker supervisor started') def started(): self._update_lifecycle('running') if startup_file: write_startup_result( startup_file, launch_nonce, 'ready', instance_id=self.instance_id, pid=self.identity.pid, ) self._startup_ready = True def interrupted(signum, _frame): self.request_interrupt(signum) if threading.current_thread() is threading.main_thread(): for name in ('SIGINT', 'SIGTERM'): current = getattr(signal, name, None) if current is not None: previous_handlers[current] = signal.getsignal(current) signal.signal(current, interrupted) client_exit = run_client( self.args, drain_event=self.drain_event, stop_event=self.stop_event, event_callback=self._event, terminal_callback=self._terminal, log_callback=self.local.log, acquire_lock=False, started_callback=started, diagnostic_callback=self._diagnostic, progress_event_reader=self.local.events_after, ) exit_code = int(client_exit or 0) if self._deadline_escalated and exit_code == 0: exit_code = 2 except BaseException as exc: if self.local is not None: self.local.log(f'worker supervisor failed: {type(exc).__name__}') if startup_file and not self._startup_ready: write_startup_result(startup_file, launch_nonce, 'failed', error=type(exc).__name__) if isinstance(exc, (KeyboardInterrupt, SystemExit)): exit_code = int(getattr(exc, 'code', 1) or 0) elif isinstance(exc, WorkerAlreadyRunning): raise finally: self._stop_maintenance() for current, previous in previous_handlers.items(): signal.signal(current, previous) try: if self.record is not None: try: pending_slots = persisted_slot_ids(self.state_dir) except OSError: pending_slots = {-1} for slot in self.local.snapshot()['slots']: if slot['phase'] != WorkerPhase.STOPPED.value: if slot['phase'] != WorkerPhase.DRAINING.value: try: self.local.emit_phase( self.instance_id, slot['slot_id'], WorkerPhase.DRAINING, reservation_id=slot['reservation_id'], source=slot['source'], scan_deadline_at=slot['scan_deadline_at'], assignment_deadline_at=slot['assignment_deadline_at'], progress={'reason': 'supervisor_shutdown'}, ) except ValueError: pass if slot['slot_id'] not in pending_slots and not self._deadline_escalated: try: self.local.emit_phase( self.instance_id, slot['slot_id'], WorkerPhase.STOPPED, progress={'reason': 'supervisor_shutdown'}, ) except ValueError: pass if self._deadline_escalated: forced_code = int(exit_code or 2) if forced_code == 0: forced_code = 2 self._hard_exit_invoked = True self._write_receipt(forced_code, drained=False) self._hard_exit_hook(forced_code) return forced_code try: self._write_receipt(exit_code) finally: self._remove_instance() self._shutdown_control() if not self._deadline_escalated: try: self.local.cleanup_retention({ 'bundles': self.args.bundle_dir, 'work': self.args.work_dir, }) except Exception as exc: self.local.log(f'worker retention cleanup failed: {type(exc).__name__}') finally: if not self._hard_exit_invoked: self._lock.release() return exit_code def write_startup_result(path, launch_nonce, outcome, *, instance_id=None, pid=None, error=None): value = { 'schema': STARTUP_SCHEMA, 'launch_nonce': str(launch_nonce or ''), 'outcome': str(outcome), 'instance_id': str(instance_id) if instance_id else None, 'pid': int(pid) if pid else None, 'error': str(error) if error else None, 'completed_at': utc_now(), } if value['outcome'] not in {'ready', 'already_running', 'failed'}: raise WorkerSupervisorError('startup result outcome is invalid') atomic_write_private_json(path, value) return value def load_startup_result(path, launch_nonce): value = read_private_json(path) if not isinstance(value, dict) or set(value) != { 'schema', 'launch_nonce', 'outcome', 'instance_id', 'pid', 'error', 'completed_at', } or value.get('schema') != STARTUP_SCHEMA: raise WorkerSupervisorError('startup result shape is invalid') if not hmac.compare_digest(str(value.get('launch_nonce') or ''), str(launch_nonce)): raise WorkerSupervisorError('startup result nonce mismatch') if value.get('outcome') not in {'ready', 'already_running', 'failed'}: raise WorkerSupervisorError('startup result outcome is invalid') if not isinstance(value.get('completed_at'), str) or not value['completed_at'].endswith('Z'): raise WorkerSupervisorError('startup result timestamp is invalid') if value['outcome'] == 'ready': if ( not isinstance(value.get('instance_id'), str) or not value['instance_id'] or type(value.get('pid')) is not int or value['pid'] <= 0 or value.get('error') is not None ): raise WorkerSupervisorError('ready startup result is invalid') elif value.get('instance_id') is not None or value.get('pid') is not None: raise WorkerSupervisorError('failed startup result is invalid') if value.get('error') is not None and not isinstance(value['error'], str): raise WorkerSupervisorError('startup result error is invalid') return value def detached_command(bootstrap_path, launch_file, startup_file, launch_nonce): return [ sys.executable, '-u', '-I', '-S', '-B', os.path.abspath(bootstrap_path), '--', '_supervise', '--launch-file', os.path.abspath(launch_file), '--startup-file', os.path.abspath(startup_file), '--launch-nonce', str(launch_nonce), ] def spawn_detached(command, *, popen=subprocess.Popen, platform_name=None): platform_name = os.name if platform_name is None else platform_name options = { 'stdin': subprocess.DEVNULL, 'stdout': subprocess.DEVNULL, 'stderr': subprocess.DEVNULL, 'close_fds': True, } if platform_name == 'nt': options['creationflags'] = ( getattr(subprocess, 'CREATE_NEW_PROCESS_GROUP', 0) | getattr(subprocess, 'DETACHED_PROCESS', 0) ) else: options['start_new_session'] = True return popen(command, **options) def wait_for_startup(startup_file, launch_nonce, process, timeout=30.0): deadline = time.monotonic() + max(0.1, float(timeout)) while time.monotonic() < deadline: if os.path.exists(startup_file): return load_startup_result(startup_file, launch_nonce) if process.poll() is not None: raise WorkerSupervisorError('worker supervisor exited before startup handshake') time.sleep(0.05) raise WorkerSupervisorError('worker supervisor startup handshake timed out') def load_shutdown_receipt(state_dir, expected_instance_id=None): value = read_private_json(shutdown_path(state_dir)) if not isinstance(value, dict) or set(value) != { 'schema', 'instance_id', 'completed_at', 'exit_code', 'drained', 'last_sequence', } or value.get('schema') != SHUTDOWN_SCHEMA: raise WorkerSupervisorError('shutdown receipt shape is invalid') if expected_instance_id is not None and not hmac.compare_digest( str(value.get('instance_id') or ''), str(expected_instance_id), ): raise WorkerSupervisorError('shutdown receipt instance mismatch') if ( not isinstance(value.get('instance_id'), str) or not value['instance_id'] or not isinstance(value.get('completed_at'), str) or not value['completed_at'].endswith('Z') or type(value.get('exit_code')) is not int or type(value.get('drained')) is not bool or type(value.get('last_sequence')) is not int or value['last_sequence'] < 0 ): raise WorkerSupervisorError('shutdown receipt fields are invalid') return value