1104 lines
45 KiB
Python
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
|