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

1104 lines
45 KiB
Python

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