Files
neuron-tai/packages/node/meshnet_node/native_worker_supervisor.py

417 lines
16 KiB
Python

"""Lifecycle supervision for the standalone native Shard worker (DGR-040).
This module deliberately has no dependency on ``TorchNodeServer``. A native
worker is an optional backend process; a failed worker must withdraw only its
own capability, never mutate or stop the existing Transformers backend. DGR-041
will connect the availability callbacks to backend-agnostic registration.
"""
from __future__ import annotations
import hashlib
import os
import re
import signal
import subprocess
import threading
from collections import deque
from collections.abc import Callable, Mapping
from dataclasses import dataclass, field
from pathlib import Path
from .native_protocol import SCHEMA_VERSION, pb
class NativeWorkerError(RuntimeError):
"""The configured worker cannot safely be started or trusted."""
_SHA256 = re.compile(r"^[0-9a-f]{64}$")
@dataclass(frozen=True)
class NativeWorkerSpec:
"""The immutable identity and launch command for one native worker."""
binary: Path
binary_digest: str
listen_address: str
artifact_path: Path
artifact_digest: str
recipe_digest: str
recipe_id: str
recipe_version: str
catalogue_version: str
shard_start: int
shard_end: int
args: tuple[str, ...] = ()
extra_environment: Mapping[str, str] = field(default_factory=dict)
def __post_init__(self) -> None:
if not self.listen_address:
raise ValueError("native worker requires a listen address")
if self.shard_start < 0 or self.shard_end <= self.shard_start:
raise ValueError("native worker range must be a non-empty half-open range")
for name in ("binary_digest", "artifact_digest", "recipe_digest"):
if not _SHA256.fullmatch(getattr(self, name)):
raise ValueError(f"native worker requires a lowercase SHA-256 {name}")
for name in ("recipe_id", "recipe_version", "catalogue_version"):
if not getattr(self, name):
raise ValueError(f"native worker requires {name}")
def environment(self) -> dict[str, str]:
"""Return the one startup identity the C++ worker must receive."""
result = dict(os.environ)
result.update({str(key): str(value) for key, value in self.extra_environment.items()})
result.update(
{
"MESHNET_SHARD_LISTEN_ADDR": self.listen_address,
"MESHNET_MODEL_ARTIFACT": str(self.artifact_path),
"MESHNET_MODEL_ARTIFACT_DIGEST": self.artifact_digest,
"MESHNET_RUNTIME_RECIPE_DIGEST": self.recipe_digest,
"MESHNET_RECIPE_ID": self.recipe_id,
"MESHNET_RECIPE_VERSION": self.recipe_version,
"MESHNET_CATALOGUE_VERSION": self.catalogue_version,
"MESHNET_SHARD_START_LAYER": str(self.shard_start),
"MESHNET_SHARD_END_LAYER": str(self.shard_end),
}
)
return result
@dataclass(frozen=True)
class NativeWorkerProbe:
"""The capability/health facts accepted by supervision after process launch."""
artifact_digest: str
recipe_digest: str
recipe_id: str
recipe_version: str
catalogue_version: str
shard_start: int
shard_end: int
serving: bool
detail: str = ""
WorkerProbe = Callable[[NativeWorkerSpec, float], NativeWorkerProbe]
AvailabilityCallback = Callable[[str], None]
class NativeWorkerSupervisor:
"""Own one worker process, its bounded logs, readiness and availability.
``start`` does not make a capability available merely because a child was
spawned: it verifies the executable and artifact bytes, waits for the
worker's readiness line, then proves the worker's reported identity and
serving health. A caller may inject ``probe`` for model-free tests; the
default performs the real gRPC capability and health calls.
"""
def __init__(
self,
spec: NativeWorkerSpec,
*,
probe: WorkerProbe | None = None,
readiness_timeout: float = 15.0,
health_timeout: float = 3.0,
health_interval: float = 5.0,
shutdown_timeout: float = 10.0,
kill_timeout: float = 3.0,
log_lines: int = 200,
on_available: AvailabilityCallback | None = None,
on_unavailable: AvailabilityCallback | None = None,
) -> None:
if min(readiness_timeout, health_timeout, health_interval, shutdown_timeout, kill_timeout) <= 0:
raise ValueError("native worker timeouts must be positive")
self.spec = spec
self._probe = probe or _grpc_probe
self._readiness_timeout = readiness_timeout
self._health_timeout = health_timeout
self._health_interval = health_interval
self._shutdown_timeout = shutdown_timeout
self._kill_timeout = kill_timeout
self._logs: deque[str] = deque(maxlen=log_lines)
self._on_available = on_available
self._on_unavailable = on_unavailable
self._process: subprocess.Popen[str] | None = None
self._ready = threading.Event()
self._stop_monitor = threading.Event()
self._lock = threading.RLock()
self._monitor: threading.Thread | None = None
self._available = False
self._unavailable_reason = "not started"
self._generation = 0
@property
def available(self) -> bool:
with self._lock:
return self._available
@property
def unavailable_reason(self) -> str:
with self._lock:
return self._unavailable_reason
@property
def logs(self) -> tuple[str, ...]:
with self._lock:
return tuple(self._logs)
@property
def pid(self) -> int | None:
with self._lock:
return None if self._process is None else self._process.pid
def start(self) -> NativeWorkerProbe:
"""Start and verify a previously stopped worker before publishing it."""
with self._lock:
if self._process is not None and self._process.poll() is None:
raise NativeWorkerError("native worker is already running; use restart()")
self._verify_startup_inputs()
self._ready.clear()
self._stop_monitor.clear()
command = [str(self.spec.binary), *self.spec.args]
try:
self._process = subprocess.Popen(
command,
stdin=subprocess.DEVNULL,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
bufsize=1,
env=self.spec.environment(),
start_new_session=True,
)
except OSError as exc:
self._process = None
raise NativeWorkerError(f"could not start native worker: {exc}") from exc
self._generation += 1
generation = self._generation
process = self._process
for stream_name, stream in (("stdout", process.stdout), ("stderr", process.stderr)):
assert stream is not None
threading.Thread(
target=self._capture_stream,
args=(stream_name, stream),
daemon=True,
).start()
if not self._ready.wait(self._readiness_timeout):
self._fail_start("worker did not report readiness before timeout")
if process.poll() is not None:
self._fail_start(f"worker exited during startup with code {process.returncode}")
try:
result = self._probe(self.spec, self._health_timeout)
self._verify_probe(result)
except Exception as exc:
self._fail_start(f"worker failed capability/health probe: {exc}")
with self._lock:
if self._process is not process or process.poll() is not None:
self._fail_start("worker exited while capability was being verified")
self._available = True
self._unavailable_reason = ""
self._monitor = threading.Thread(
target=self._monitor_loop, args=(generation, process), daemon=True
)
self._monitor.start()
if self._on_available is not None:
self._on_available("worker ready and identity verified")
return result
def add_availability_callbacks(
self,
*,
on_available: AvailabilityCallback | None = None,
on_unavailable: AvailabilityCallback | None = None,
) -> None:
"""Attach an integration callback before the worker is started.
Registration is deliberately supplied by the caller so this supervisor
stays independent of Tracker HTTP and of every other backend.
"""
with self._lock:
if self._process is not None:
raise NativeWorkerError("availability callbacks must be attached before start")
self._on_available = _combine_callbacks(self._on_available, on_available)
self._on_unavailable = _combine_callbacks(self._on_unavailable, on_unavailable)
def restart(self) -> NativeWorkerProbe:
"""Withdraw the old capability, stop its process, then prove a fresh one."""
self.stop(reason="worker restart requested")
return self.start()
def stop(self, *, reason: str = "worker stopped") -> None:
"""Gracefully terminate the owned process, escalating only after a bound."""
with self._lock:
process = self._process
self._stop_monitor.set()
self._process = None
self._mark_unavailable(reason)
if process is None or process.poll() is not None:
return
_terminate_process_group(process, self._shutdown_timeout, self._kill_timeout)
def check_health(self) -> bool:
"""Run one bounded health check and withdraw availability on failure."""
with self._lock:
process = self._process
if process is None or process.poll() is not None:
self._mark_unavailable("worker process exited")
return False
try:
result = self._probe(self.spec, self._health_timeout)
self._verify_probe(result)
except Exception as exc:
self._mark_unavailable(f"worker health lost: {exc}")
return False
return True
def _verify_startup_inputs(self) -> None:
if not self.spec.binary.is_file() or not os.access(self.spec.binary, os.X_OK):
raise NativeWorkerError(f"native worker binary is not executable: {self.spec.binary}")
if _sha256_file(self.spec.binary) != self.spec.binary_digest:
raise NativeWorkerError("native worker binary digest does not match its immutable pin")
if not self.spec.artifact_path.is_file():
raise NativeWorkerError(f"native worker artifact is missing: {self.spec.artifact_path}")
digest = _sha256_file(self.spec.artifact_path)
if digest != self.spec.artifact_digest:
raise NativeWorkerError("native worker artifact digest does not match its immutable pin")
def _verify_probe(self, probe: NativeWorkerProbe) -> None:
expected = self.spec
actual = (
probe.artifact_digest,
probe.recipe_digest,
probe.recipe_id,
probe.recipe_version,
probe.catalogue_version,
probe.shard_start,
probe.shard_end,
)
wanted = (
expected.artifact_digest,
expected.recipe_digest,
expected.recipe_id,
expected.recipe_version,
expected.catalogue_version,
expected.shard_start,
expected.shard_end,
)
if actual != wanted:
raise NativeWorkerError("worker probe identity/range differs from configured startup identity")
if not probe.serving:
raise NativeWorkerError(f"worker is not serving: {probe.detail or 'no detail'}")
def _capture_stream(self, stream_name: str, stream) -> None:
for raw_line in stream:
line = f"{stream_name}: {raw_line.rstrip()}"
with self._lock:
self._logs.append(line)
if raw_line.startswith("ShardRuntime worker listening on "):
self._ready.set()
def _monitor_loop(self, generation: int, process: subprocess.Popen[str]) -> None:
while not self._stop_monitor.wait(self._health_interval):
with self._lock:
if generation != self._generation or self._process is not process:
return
if process.poll() is not None:
self._mark_unavailable(f"worker process exited with code {process.returncode}")
return
if not self.check_health():
return
def _fail_start(self, reason: str) -> None:
self.stop(reason=reason)
raise NativeWorkerError(reason)
def _mark_unavailable(self, reason: str) -> None:
callback = None
with self._lock:
was_available = self._available
self._available = False
self._unavailable_reason = reason
if was_available:
callback = self._on_unavailable
if callback is not None:
callback(reason)
def _sha256_file(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as file:
for chunk in iter(lambda: file.read(1024 * 1024), b""):
digest.update(chunk)
return digest.hexdigest()
def _combine_callbacks(
first: AvailabilityCallback | None, second: AvailabilityCallback | None
) -> AvailabilityCallback | None:
if first is None:
return second
if second is None:
return first
def combined(reason: str) -> None:
first(reason)
second(reason)
return combined
def _terminate_process_group(
process: subprocess.Popen[str], shutdown_timeout: float, kill_timeout: float
) -> None:
try:
os.killpg(process.pid, signal.SIGTERM)
except ProcessLookupError:
return
try:
process.wait(timeout=shutdown_timeout)
return
except subprocess.TimeoutExpired:
pass
try:
os.killpg(process.pid, signal.SIGKILL)
except ProcessLookupError:
return
try:
process.wait(timeout=kill_timeout)
except subprocess.TimeoutExpired as exc:
raise NativeWorkerError("native worker did not terminate after SIGKILL") from exc
def _grpc_probe(spec: NativeWorkerSpec, timeout: float) -> NativeWorkerProbe:
"""Default real wire probe; importing grpc lazily preserves CLI startup."""
import grpc
from .native_protocol.generated import shard_runtime_pb2_grpc as pb_grpc
channel = grpc.insecure_channel(spec.listen_address)
try:
grpc.channel_ready_future(channel).result(timeout=timeout)
stub = pb_grpc.ShardRuntimeStub(channel)
capability = stub.GetCapability(pb.CapabilityRequest(schema_version=SCHEMA_VERSION), timeout=timeout)
health = stub.Health(pb.HealthRequest(schema_version=SCHEMA_VERSION), timeout=timeout)
finally:
channel.close()
fingerprint = capability.fingerprint
shard_range = capability.shard_range
return NativeWorkerProbe(
artifact_digest=fingerprint.model_artifact_digest,
recipe_digest=fingerprint.runtime_recipe_digest,
recipe_id=fingerprint.recipe_id,
recipe_version=fingerprint.recipe_version,
catalogue_version=fingerprint.catalogue_version,
shard_start=shard_range.start_layer,
shard_end=shard_range.end_layer,
serving=(
capability.validated
and health.state == pb.SERVING_STATE_SERVING
),
detail=health.detail or capability.detail,
)