Initial public ModelForge release
This commit is contained in:
@@ -0,0 +1 @@
|
||||
__version__ = "1.2.2"
|
||||
@@ -0,0 +1,372 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import shutil
|
||||
import ssl
|
||||
import struct
|
||||
from pathlib import Path, PurePosixPath
|
||||
from typing import Any, Literal
|
||||
from urllib.parse import quote
|
||||
|
||||
import certifi
|
||||
import httpx
|
||||
from modelforge_api.domain.acquisition import (
|
||||
AgentArtifactJobLease,
|
||||
AgentJobComplete,
|
||||
AgentJobFailure,
|
||||
AgentJobProgress,
|
||||
CompletedFile,
|
||||
)
|
||||
|
||||
from modelforge_node_agent.settings import AgentSettings
|
||||
|
||||
|
||||
class AcquisitionFailure(RuntimeError):
|
||||
def __init__(self, code: str, message: str, *, retryable: bool = False) -> None:
|
||||
super().__init__(message)
|
||||
self.code = code
|
||||
self.retryable = retryable
|
||||
|
||||
|
||||
class AcquisitionCancelled(AcquisitionFailure):
|
||||
def __init__(self) -> None:
|
||||
super().__init__("cancelled", "artifact acquisition was cancelled")
|
||||
|
||||
|
||||
def _confined(base: Path, candidate: Path) -> Path:
|
||||
resolved_base = base.resolve()
|
||||
resolved = candidate.resolve()
|
||||
if resolved != resolved_base and resolved_base not in resolved.parents:
|
||||
raise AcquisitionFailure("path_escape", "artifact path escaped its approved storage root")
|
||||
return resolved
|
||||
|
||||
|
||||
def _relative_path(value: str) -> PurePosixPath:
|
||||
path = PurePosixPath(value.replace("\\", "/"))
|
||||
if path.is_absolute() or not path.parts or ".." in path.parts:
|
||||
raise AcquisitionFailure("invalid_path", "upstream file path is not confined")
|
||||
return path
|
||||
|
||||
|
||||
def _hash_file(path: Path) -> tuple[str, int]:
|
||||
digest = hashlib.sha256()
|
||||
size = 0
|
||||
with path.open("rb") as handle:
|
||||
for chunk in iter(lambda: handle.read(1024 * 1024), b""):
|
||||
digest.update(chunk)
|
||||
size += len(chunk)
|
||||
return digest.hexdigest(), size
|
||||
|
||||
|
||||
def _inspect_file(path: Path, plan_file: Any) -> list[dict[str, Any]]:
|
||||
findings: list[dict[str, Any]] = []
|
||||
if path.is_symlink():
|
||||
return [{"type": "symlink", "status": "blocked", "severity": "block", "evidence": {}}]
|
||||
risks = set(plan_file.risk_flags)
|
||||
if "pickle_or_executable_serialization" in risks:
|
||||
findings.append(
|
||||
{
|
||||
"type": "serialization",
|
||||
"status": "blocked",
|
||||
"severity": "block",
|
||||
"evidence": {"format": plan_file.file_format},
|
||||
}
|
||||
)
|
||||
if plan_file.file_format == "python" or "remote_code" in risks:
|
||||
findings.append(
|
||||
{
|
||||
"type": "remote_code",
|
||||
"status": "blocked",
|
||||
"severity": "block",
|
||||
"evidence": {"trust_remote_code": False},
|
||||
}
|
||||
)
|
||||
if plan_file.file_format == "safetensors":
|
||||
with path.open("rb") as handle:
|
||||
raw = handle.read(8)
|
||||
if len(raw) != 8:
|
||||
raise AcquisitionFailure(
|
||||
"invalid_safetensors", f"truncated safetensors header: {plan_file.path}"
|
||||
)
|
||||
header_size = struct.unpack("<Q", raw)[0]
|
||||
if header_size > 100 * 1024 * 1024 or header_size > path.stat().st_size - 8:
|
||||
raise AcquisitionFailure(
|
||||
"invalid_safetensors", f"unsafe safetensors header: {plan_file.path}"
|
||||
)
|
||||
try:
|
||||
header = json.loads(handle.read(header_size))
|
||||
except (UnicodeDecodeError, json.JSONDecodeError) as exc:
|
||||
raise AcquisitionFailure(
|
||||
"invalid_safetensors", f"invalid safetensors JSON header: {plan_file.path}"
|
||||
) from exc
|
||||
if not isinstance(header, dict):
|
||||
raise AcquisitionFailure(
|
||||
"invalid_safetensors", f"invalid safetensors structure: {plan_file.path}"
|
||||
)
|
||||
findings.append(
|
||||
{
|
||||
"type": "safetensors_header",
|
||||
"status": "passed",
|
||||
"severity": "info",
|
||||
"evidence": {"header_bytes": header_size, "tensor_entries": len(header)},
|
||||
}
|
||||
)
|
||||
if plan_file.file_format == "json" and path.stat().st_size <= 16 * 1024 * 1024:
|
||||
try:
|
||||
payload = json.loads(path.read_text(encoding="utf-8"))
|
||||
except (UnicodeDecodeError, json.JSONDecodeError):
|
||||
payload = None
|
||||
if isinstance(payload, dict) and payload.get("auto_map"):
|
||||
findings.append(
|
||||
{
|
||||
"type": "remote_code_configuration",
|
||||
"status": "blocked",
|
||||
"severity": "block",
|
||||
"evidence": {"auto_map_present": True, "trust_remote_code": False},
|
||||
}
|
||||
)
|
||||
if not findings:
|
||||
findings.append(
|
||||
{
|
||||
"type": "file_policy",
|
||||
"status": "passed",
|
||||
"severity": "info",
|
||||
"evidence": {"format": plan_file.file_format},
|
||||
}
|
||||
)
|
||||
return findings
|
||||
|
||||
|
||||
class ArtifactAcquirer:
|
||||
def __init__(self, settings: AgentSettings, transport: Any, credential: str) -> None:
|
||||
self.settings = settings
|
||||
self.transport = transport
|
||||
self.credential = credential
|
||||
|
||||
def _progress(
|
||||
self,
|
||||
lease: AgentArtifactJobLease,
|
||||
status: Literal["claimed", "downloading", "verifying", "promoting"],
|
||||
progress_bytes: int,
|
||||
current_file: str | None,
|
||||
quarantine_relative_path: str | None,
|
||||
) -> None:
|
||||
response = self.transport.artifact_progress(
|
||||
str(lease.job_id),
|
||||
AgentJobProgress(
|
||||
lease_token=lease.lease_token,
|
||||
status=status,
|
||||
progress_bytes=progress_bytes,
|
||||
current_file=current_file,
|
||||
quarantine_relative_path=quarantine_relative_path,
|
||||
).model_dump(mode="json"),
|
||||
self.credential,
|
||||
)
|
||||
if response.get("cancel_requested"):
|
||||
raise AcquisitionCancelled()
|
||||
|
||||
def _download(
|
||||
self,
|
||||
client: httpx.Client,
|
||||
lease: AgentArtifactJobLease,
|
||||
plan_file: Any,
|
||||
destination: Path,
|
||||
progress_before: int,
|
||||
quarantine_relative_path: str,
|
||||
) -> None:
|
||||
destination.parent.mkdir(parents=True, exist_ok=True)
|
||||
partial = destination.with_suffix(destination.suffix + ".part")
|
||||
offset = partial.stat().st_size if partial.is_file() else 0
|
||||
if offset > plan_file.size_bytes:
|
||||
partial.unlink()
|
||||
offset = 0
|
||||
elif offset == plan_file.size_bytes:
|
||||
# A transport may fail while closing a response after every expected byte was
|
||||
# durably written. Treat the complete partial as resumable input and let the
|
||||
# normal verification phase decide whether its digest is acceptable.
|
||||
os.replace(partial, destination)
|
||||
return
|
||||
headers: dict[str, str] = {}
|
||||
if offset:
|
||||
headers["Range"] = f"bytes={offset}-"
|
||||
token = self.settings.hf_token
|
||||
if token and token.get_secret_value():
|
||||
headers["Authorization"] = f"Bearer {token.get_secret_value()}"
|
||||
url = (
|
||||
"https://huggingface.co/"
|
||||
f"{quote(lease.repository_id, safe='/')}/resolve/"
|
||||
f"{quote(lease.resolved_commit_sha, safe='')}/{quote(plan_file.path, safe='/')}"
|
||||
)
|
||||
with client.stream("GET", url, headers=headers) as response:
|
||||
response.raise_for_status()
|
||||
append = offset > 0 and response.status_code == 206
|
||||
if not append:
|
||||
offset = 0
|
||||
mode = "ab" if append else "wb"
|
||||
last_report = offset
|
||||
with partial.open(mode) as handle:
|
||||
for chunk in response.iter_bytes(1024 * 1024):
|
||||
if offset + len(chunk) > plan_file.size_bytes:
|
||||
raise AcquisitionFailure(
|
||||
"download_size_mismatch",
|
||||
f"download exceeded declared size for {plan_file.path}",
|
||||
retryable=False,
|
||||
)
|
||||
handle.write(chunk)
|
||||
offset += len(chunk)
|
||||
if offset - last_report >= 8 * 1024 * 1024:
|
||||
self._progress(
|
||||
lease,
|
||||
"downloading",
|
||||
progress_before + offset,
|
||||
plan_file.path,
|
||||
quarantine_relative_path,
|
||||
)
|
||||
last_report = offset
|
||||
handle.flush()
|
||||
os.fsync(handle.fileno())
|
||||
if offset != plan_file.size_bytes:
|
||||
raise AcquisitionFailure(
|
||||
"download_size_mismatch",
|
||||
f"downloaded size mismatch for {plan_file.path}",
|
||||
retryable=True,
|
||||
)
|
||||
os.replace(partial, destination)
|
||||
|
||||
def execute(self, raw_lease: dict[str, Any]) -> None:
|
||||
lease = AgentArtifactJobLease.model_validate(raw_lease)
|
||||
configured_root = self.settings.artifact_path.resolve()
|
||||
root = _confined(configured_root, Path(lease.target_root))
|
||||
root.mkdir(parents=True, exist_ok=True)
|
||||
if root.is_symlink():
|
||||
raise AcquisitionFailure("symlink_root", "approved artifact root is a symlink")
|
||||
usage = shutil.disk_usage(root)
|
||||
reserve = max(lease.reserve_bytes, usage.total * lease.reserve_percent // 100)
|
||||
if lease.total_size_bytes > max(0, usage.free - reserve):
|
||||
raise AcquisitionFailure("capacity_changed", "execution capacity preflight failed")
|
||||
safe_repo = lease.repository_id.replace("/", "--").replace("\\", "--")
|
||||
final = _confined(root, root / "repositories" / safe_repo / lease.resolved_commit_sha)
|
||||
stage = _confined(root, root / ".quarantine" / str(lease.job_id))
|
||||
reconcile_promoted = final.exists()
|
||||
if reconcile_promoted:
|
||||
existing_manifest = final / ".modelforge-manifest.json"
|
||||
if not existing_manifest.is_file():
|
||||
raise AcquisitionFailure(
|
||||
"promotion_conflict", "existing target has no ModelForge manifest"
|
||||
)
|
||||
existing = json.loads(existing_manifest.read_text(encoding="utf-8"))
|
||||
if (
|
||||
existing.get("repository_id") != lease.repository_id
|
||||
or existing.get("resolved_commit_sha") != lease.resolved_commit_sha
|
||||
):
|
||||
raise AcquisitionFailure("promotion_conflict", "existing target provenance differs")
|
||||
else:
|
||||
stage.mkdir(parents=True, exist_ok=True)
|
||||
worktree = final if reconcile_promoted else stage
|
||||
quarantine_relative = stage.relative_to(root).as_posix()
|
||||
self._progress(lease, "claimed", 0, None, quarantine_relative)
|
||||
completed: list[CompletedFile] = []
|
||||
progress = 0
|
||||
with httpx.Client(
|
||||
timeout=self.settings.artifact_download_timeout_seconds,
|
||||
follow_redirects=True,
|
||||
verify=ssl.create_default_context(cafile=certifi.where()),
|
||||
) as client:
|
||||
for plan_file in lease.files:
|
||||
relative = _relative_path(plan_file.path)
|
||||
destination = _confined(worktree, worktree.joinpath(*relative.parts))
|
||||
if not destination.is_file():
|
||||
if reconcile_promoted:
|
||||
raise AcquisitionFailure(
|
||||
"promotion_conflict",
|
||||
f"promoted set is incomplete: {plan_file.path}",
|
||||
)
|
||||
self._download(
|
||||
client,
|
||||
lease,
|
||||
plan_file,
|
||||
destination,
|
||||
progress,
|
||||
quarantine_relative,
|
||||
)
|
||||
self._progress(
|
||||
lease,
|
||||
"verifying",
|
||||
progress + plan_file.size_bytes,
|
||||
plan_file.path,
|
||||
quarantine_relative,
|
||||
)
|
||||
digest, size = _hash_file(destination)
|
||||
if size != plan_file.size_bytes:
|
||||
raise AcquisitionFailure(
|
||||
"verification_size_mismatch", f"size mismatch for {plan_file.path}"
|
||||
)
|
||||
if plan_file.upstream_sha256 and digest != plan_file.upstream_sha256:
|
||||
raise AcquisitionFailure(
|
||||
"verification_digest_mismatch", f"digest mismatch for {plan_file.path}"
|
||||
)
|
||||
inspections = _inspect_file(destination, plan_file)
|
||||
if any(item["severity"] == "block" for item in inspections):
|
||||
raise AcquisitionFailure(
|
||||
"static_inspection_blocked", f"static inspection blocked {plan_file.path}"
|
||||
)
|
||||
progress += size
|
||||
completed.append(
|
||||
CompletedFile(
|
||||
path=plan_file.path,
|
||||
relative_path="pending",
|
||||
size_bytes=size,
|
||||
sha256=digest,
|
||||
inspections=inspections,
|
||||
)
|
||||
)
|
||||
self._progress(lease, "promoting", progress, None, quarantine_relative)
|
||||
if not reconcile_promoted:
|
||||
final.parent.mkdir(parents=True, exist_ok=True)
|
||||
manifest = {
|
||||
"job_id": str(lease.job_id),
|
||||
"repository_id": lease.repository_id,
|
||||
"resolved_commit_sha": lease.resolved_commit_sha,
|
||||
"files": [item.model_dump(mode="json") for item in completed],
|
||||
}
|
||||
(stage / ".modelforge-manifest.json").write_text(
|
||||
json.dumps(manifest, sort_keys=True, indent=2), encoding="utf-8"
|
||||
)
|
||||
os.replace(stage, final)
|
||||
promoted_base = final.relative_to(root).as_posix()
|
||||
promoted_files = [
|
||||
item.model_copy(update={"relative_path": f"{promoted_base}/{item.path}"})
|
||||
for item in completed
|
||||
]
|
||||
after = shutil.disk_usage(root)
|
||||
self.transport.artifact_complete(
|
||||
str(lease.job_id),
|
||||
AgentJobComplete(
|
||||
lease_token=lease.lease_token,
|
||||
promoted_relative_path=promoted_base,
|
||||
capacity_observation={
|
||||
"capacity_bytes": after.total,
|
||||
"free_bytes": after.free,
|
||||
"reserve_bytes": reserve,
|
||||
"agent_path": str(root),
|
||||
},
|
||||
files=promoted_files,
|
||||
).model_dump(mode="json"),
|
||||
self.credential,
|
||||
)
|
||||
|
||||
def fail(self, raw_lease: dict[str, Any], error: AcquisitionFailure) -> None:
|
||||
lease = AgentArtifactJobLease.model_validate(raw_lease)
|
||||
self.transport.artifact_fail(
|
||||
str(lease.job_id),
|
||||
AgentJobFailure(
|
||||
lease_token=lease.lease_token,
|
||||
error_code=error.code,
|
||||
error_message=str(error),
|
||||
retryable=error.retryable,
|
||||
details={"exception_type": type(error).__name__},
|
||||
).model_dump(mode="json"),
|
||||
self.credential,
|
||||
)
|
||||
@@ -0,0 +1,309 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
from datetime import UTC, datetime
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from modelforge_api.domain.agent_protocol import (
|
||||
AGENT_PROTOCOL_CAPABILITIES,
|
||||
AGENT_PROTOCOL_VERSION,
|
||||
AgentMetadata,
|
||||
EnrollmentRequest,
|
||||
HeartbeatRequest,
|
||||
InventoryNvidiaPayload,
|
||||
InventoryReport,
|
||||
TelemetryReport,
|
||||
)
|
||||
from modelforge_api.domain.hardware import NvidiaCollection
|
||||
from modelforge_api.hardware.collectors import (
|
||||
NodeIdentityProvider,
|
||||
SystemHostCollector,
|
||||
build_nvml_collector,
|
||||
)
|
||||
|
||||
from modelforge_node_agent import __version__
|
||||
from modelforge_node_agent.acquisition import AcquisitionFailure, ArtifactAcquirer
|
||||
from modelforge_node_agent.preflight import (
|
||||
AgentObservationError,
|
||||
AgentStartupError,
|
||||
nvidia_required,
|
||||
validate_nvidia_collection,
|
||||
)
|
||||
from modelforge_node_agent.settings import AgentSettings
|
||||
from modelforge_node_agent.state import AgentPersistentState, AgentStateStore
|
||||
from modelforge_node_agent.transport import AgentTransport
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def describe_failure(error: BaseException) -> str:
|
||||
"""A bounded, non-secret description an operator can act on.
|
||||
|
||||
`type(error).__name__` alone cannot distinguish a revoked credential from a DNS failure, which
|
||||
is exactly the question an operator asks when an agent stops publishing. The status code, the
|
||||
method and the path answer it. The response body and the credential never appear: the body can
|
||||
echo request content and the credential is the secret itself.
|
||||
"""
|
||||
|
||||
if isinstance(error, AgentObservationError):
|
||||
return str(error.problem)
|
||||
if isinstance(error, httpx.HTTPStatusError):
|
||||
request = error.request
|
||||
return (
|
||||
f"{type(error).__name__} {error.response.status_code} "
|
||||
f"on {request.method} {request.url.path}"
|
||||
)
|
||||
if isinstance(error, httpx.RequestError):
|
||||
# A transport error raised before the request was attached has no request to describe.
|
||||
try:
|
||||
request = error.request
|
||||
except RuntimeError:
|
||||
return type(error).__name__
|
||||
return f"{type(error).__name__} on {request.method} {request.url.path}"
|
||||
return type(error).__name__
|
||||
|
||||
|
||||
class NodeAgent:
|
||||
def __init__(
|
||||
self,
|
||||
settings: AgentSettings,
|
||||
transport: Any = None,
|
||||
host_collector: Any = None,
|
||||
accelerator_collector: Any = None,
|
||||
state_store: AgentStateStore | None = None,
|
||||
) -> None:
|
||||
self.settings = settings
|
||||
self.transport = transport or AgentTransport(
|
||||
settings.control_plane_url,
|
||||
settings.request_timeout_seconds,
|
||||
settings.tls_verify,
|
||||
)
|
||||
identity = NodeIdentityProvider(
|
||||
settings.identity_file,
|
||||
settings.identity,
|
||||
force_persisted=settings.identity_mode == "persisted",
|
||||
)
|
||||
self.host_collector = host_collector or SystemHostCollector(
|
||||
identity,
|
||||
{
|
||||
"model_cache": settings.model_cache_path,
|
||||
"artifacts": settings.artifact_path,
|
||||
"quarantine": settings.quarantine_path,
|
||||
},
|
||||
)
|
||||
self.accelerator_collector = accelerator_collector or build_nvml_collector()
|
||||
self.state_store = state_store or AgentStateStore(
|
||||
settings.state_file, settings.credential_file
|
||||
)
|
||||
self.state: AgentPersistentState = self.state_store.load()
|
||||
self.credential = self.state_store.credential()
|
||||
self.started_at = datetime.now(UTC)
|
||||
self._stop = asyncio.Event()
|
||||
# The publication loop and the artifact loop run in separate worker threads and both
|
||||
# need a credential. Without this lock they can enrol concurrently and burn one
|
||||
# single-use token on two node identities for the same hardware.
|
||||
self._enrollment_lock = threading.Lock()
|
||||
self._preflight_passed = False
|
||||
|
||||
def collect_accelerators(self, *, startup: bool = False) -> NvidiaCollection:
|
||||
collection: NvidiaCollection = self.accelerator_collector.collect()
|
||||
validate_nvidia_collection(
|
||||
collection,
|
||||
required=nvidia_required(self.settings.accelerator_mode),
|
||||
startup=startup,
|
||||
)
|
||||
return collection
|
||||
|
||||
def preflight(self) -> NvidiaCollection:
|
||||
"""Validate the accelerator contract before enrollment or publication."""
|
||||
|
||||
collection = self.collect_accelerators(startup=True)
|
||||
self._preflight_passed = True
|
||||
return collection
|
||||
|
||||
def metadata(self) -> AgentMetadata:
|
||||
return AgentMetadata(
|
||||
agent_version=__version__,
|
||||
protocol_version=AGENT_PROTOCOL_VERSION,
|
||||
supported_capabilities=AGENT_PROTOCOL_CAPABILITIES,
|
||||
started_at=self.started_at,
|
||||
)
|
||||
|
||||
def ensure_enrolled(self) -> None:
|
||||
if not self._preflight_passed:
|
||||
self.preflight()
|
||||
if self.credential:
|
||||
return
|
||||
with self._enrollment_lock:
|
||||
if self.credential:
|
||||
return
|
||||
token = self.settings.enrollment_token
|
||||
if token is None or not token.get_secret_value():
|
||||
raise RuntimeError("agent is not enrolled and no enrollment token was provided")
|
||||
host = self.host_collector.collect()
|
||||
response = self.transport.enroll(
|
||||
EnrollmentRequest(
|
||||
enrollment_token=token.get_secret_value(),
|
||||
identity_key=host.identity_key,
|
||||
identity_source=host.identity_source,
|
||||
hostname=host.hostname,
|
||||
display_name=self.settings.display_name or host.display_name,
|
||||
metadata=self.metadata(),
|
||||
).model_dump(mode="json")
|
||||
)
|
||||
credential = str(response["node_credential"])
|
||||
self.state_store.save_credential(credential)
|
||||
self.credential = credential
|
||||
|
||||
def enrolled_credential(self) -> str:
|
||||
self.ensure_enrolled()
|
||||
if self.credential is None: # Defensive invariant for custom state stores.
|
||||
raise RuntimeError("enrollment did not produce a node credential")
|
||||
return self.credential
|
||||
|
||||
def heartbeat_once(self, last_error: str | None = None) -> None:
|
||||
credential = self.enrolled_credential()
|
||||
host = self.host_collector.collect()
|
||||
request = HeartbeatRequest(
|
||||
identity_key=host.identity_key,
|
||||
metadata=self.metadata(),
|
||||
observed_at=datetime.now(UTC),
|
||||
last_error=last_error,
|
||||
)
|
||||
self.transport.heartbeat(request.model_dump(mode="json"), credential)
|
||||
|
||||
def inventory_once(self) -> None:
|
||||
nvidia = self.collect_accelerators() if self._preflight_passed else self.preflight()
|
||||
credential = self.enrolled_credential()
|
||||
host = self.host_collector.collect()
|
||||
self.state.inventory_sequence += 1
|
||||
self.state_store.save(self.state)
|
||||
request = InventoryReport(
|
||||
identity_key=host.identity_key,
|
||||
protocol_version=AGENT_PROTOCOL_VERSION,
|
||||
sequence=self.state.inventory_sequence,
|
||||
observed_at=max(host.inventory_at, nvidia.observed_at),
|
||||
host=host,
|
||||
nvidia=InventoryNvidiaPayload(
|
||||
availability=nvidia.availability,
|
||||
reason=nvidia.reason,
|
||||
inventory=nvidia.inventory,
|
||||
),
|
||||
)
|
||||
self.transport.inventory(request.model_dump(mode="json"), credential)
|
||||
|
||||
def telemetry_once(self) -> None:
|
||||
nvidia = self.collect_accelerators() if self._preflight_passed else self.preflight()
|
||||
credential = self.enrolled_credential()
|
||||
host = self.host_collector.collect()
|
||||
self.state.telemetry_sequence += 1
|
||||
self.state_store.save(self.state)
|
||||
request = TelemetryReport(
|
||||
identity_key=host.identity_key,
|
||||
protocol_version=AGENT_PROTOCOL_VERSION,
|
||||
sequence=self.state.telemetry_sequence,
|
||||
observed_at=max(host.inventory_at, nvidia.observed_at),
|
||||
available_ram_bytes=host.available_ram_bytes,
|
||||
storage=host.storage,
|
||||
accelerators=nvidia.telemetry,
|
||||
)
|
||||
self.transport.telemetry(request.model_dump(mode="json"), credential)
|
||||
|
||||
async def run(self) -> None:
|
||||
if not self._preflight_passed:
|
||||
self.preflight()
|
||||
artifact_task = asyncio.create_task(self._artifact_loop())
|
||||
next_heartbeat = next_inventory = next_telemetry = 0.0
|
||||
backoff = 1
|
||||
last_error: str | None = None
|
||||
try:
|
||||
while not self._stop.is_set():
|
||||
now = time.monotonic()
|
||||
try:
|
||||
if now >= next_heartbeat:
|
||||
await asyncio.to_thread(self.heartbeat_once, last_error)
|
||||
next_heartbeat = now + self.settings.heartbeat_interval_seconds
|
||||
if now >= next_inventory:
|
||||
await asyncio.to_thread(self.inventory_once)
|
||||
next_inventory = now + self.settings.inventory_interval_seconds
|
||||
if now >= next_telemetry:
|
||||
await asyncio.to_thread(self.telemetry_once)
|
||||
next_telemetry = now + self.settings.telemetry_interval_seconds
|
||||
backoff = 1
|
||||
last_error = None
|
||||
delay = (
|
||||
min(
|
||||
next_heartbeat,
|
||||
next_inventory,
|
||||
next_telemetry,
|
||||
)
|
||||
- time.monotonic()
|
||||
)
|
||||
delay = max(0.1, delay)
|
||||
except AgentStartupError:
|
||||
raise
|
||||
except (httpx.HTTPError, OSError, RuntimeError, ValueError, KeyError) as exc:
|
||||
last_error = describe_failure(exc)
|
||||
logger.warning(
|
||||
"control-plane publication failed; retrying: %s (backoff %ss)",
|
||||
last_error,
|
||||
backoff,
|
||||
)
|
||||
delay = backoff
|
||||
backoff = min(backoff * 2, self.settings.max_backoff_seconds)
|
||||
try:
|
||||
await asyncio.wait_for(self._stop.wait(), delay)
|
||||
except TimeoutError:
|
||||
continue
|
||||
finally:
|
||||
self._stop.set()
|
||||
await artifact_task
|
||||
|
||||
def artifact_once(self) -> bool:
|
||||
credential = self.enrolled_credential()
|
||||
if not hasattr(self.transport, "next_artifact_job"):
|
||||
return False
|
||||
lease = self.transport.next_artifact_job(credential)
|
||||
if not lease:
|
||||
return False
|
||||
acquirer = ArtifactAcquirer(self.settings, self.transport, credential)
|
||||
try:
|
||||
acquirer.execute(lease)
|
||||
except AcquisitionFailure as exc:
|
||||
acquirer.fail(lease, exc)
|
||||
except httpx.HTTPError as exc:
|
||||
acquirer.fail(
|
||||
lease,
|
||||
AcquisitionFailure(
|
||||
"download_transport_error", type(exc).__name__, retryable=True
|
||||
),
|
||||
)
|
||||
except (OSError, RuntimeError, ValueError, KeyError) as exc:
|
||||
acquirer.fail(
|
||||
lease,
|
||||
AcquisitionFailure("agent_execution_error", type(exc).__name__),
|
||||
)
|
||||
return True
|
||||
|
||||
async def _artifact_loop(self) -> None:
|
||||
while not self._stop.is_set():
|
||||
try:
|
||||
worked = await asyncio.to_thread(self.artifact_once)
|
||||
delay = 0.1 if worked else self.settings.artifact_job_poll_interval_seconds
|
||||
except (httpx.HTTPError, OSError, RuntimeError, ValueError, KeyError) as exc:
|
||||
logger.warning("artifact-job poll failed; retrying: %s", describe_failure(exc))
|
||||
delay = self.settings.max_backoff_seconds
|
||||
try:
|
||||
await asyncio.wait_for(self._stop.wait(), delay)
|
||||
except TimeoutError:
|
||||
continue
|
||||
|
||||
def stop(self) -> None:
|
||||
self._stop.set()
|
||||
|
||||
def close(self) -> None:
|
||||
self.transport.close()
|
||||
@@ -0,0 +1,43 @@
|
||||
import argparse
|
||||
import asyncio
|
||||
import logging
|
||||
import signal
|
||||
|
||||
from modelforge_node_agent.agent import NodeAgent
|
||||
from modelforge_node_agent.preflight import AgentStartupError
|
||||
from modelforge_node_agent.settings import AgentSettings
|
||||
|
||||
|
||||
async def run_agent(*, preflight_only: bool = False) -> None:
|
||||
agent = NodeAgent(AgentSettings())
|
||||
loop = asyncio.get_running_loop()
|
||||
for signum in (signal.SIGINT, signal.SIGTERM):
|
||||
loop.add_signal_handler(signum, agent.stop)
|
||||
try:
|
||||
if preflight_only:
|
||||
agent.preflight()
|
||||
logging.getLogger(__name__).info("Node Agent accelerator preflight passed")
|
||||
return
|
||||
await agent.run()
|
||||
finally:
|
||||
agent.close()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description="ITWorx ModelForge Node Agent")
|
||||
parser.add_argument(
|
||||
"--preflight",
|
||||
action="store_true",
|
||||
help="validate the accelerator runtime without enrollment or publication",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
try:
|
||||
asyncio.run(run_agent(preflight_only=args.preflight))
|
||||
except AgentStartupError as exc:
|
||||
logging.getLogger(__name__).critical("%s", exc.problem)
|
||||
raise SystemExit(3) from None
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,129 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import platform
|
||||
from dataclasses import dataclass
|
||||
from enum import StrEnum
|
||||
from pathlib import Path
|
||||
|
||||
from modelforge_api.domain.enums import Availability
|
||||
from modelforge_api.domain.hardware import NvidiaCollection
|
||||
|
||||
|
||||
class AgentStartupFailureCode(StrEnum):
|
||||
NVIDIA_NVML_UNAVAILABLE = "NVIDIA_NVML_UNAVAILABLE"
|
||||
NVIDIA_DEVICE_NOT_FOUND = "NVIDIA_DEVICE_NOT_FOUND"
|
||||
NVIDIA_TELEMETRY_UNAVAILABLE = "NVIDIA_TELEMETRY_UNAVAILABLE"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AgentStartupProblem:
|
||||
code: AgentStartupFailureCode
|
||||
setting: str
|
||||
message: str
|
||||
|
||||
def __str__(self) -> str:
|
||||
return f"[{self.code}] {self.setting}: {self.message}"
|
||||
|
||||
|
||||
class AgentStartupError(RuntimeError):
|
||||
def __init__(self, problem: AgentStartupProblem) -> None:
|
||||
super().__init__(str(problem))
|
||||
self.problem = problem
|
||||
|
||||
|
||||
class AgentObservationError(RuntimeError):
|
||||
def __init__(self, problem: AgentStartupProblem) -> None:
|
||||
super().__init__(str(problem))
|
||||
self.problem = problem
|
||||
|
||||
|
||||
def nvidia_runtime_observed(
|
||||
*,
|
||||
environ: dict[str, str] | None = None,
|
||||
device_root: Path = Path("/dev"),
|
||||
) -> bool:
|
||||
"""Return whether the container has affirmative NVIDIA runtime evidence."""
|
||||
|
||||
environment = os.environ if environ is None else environ
|
||||
visible_devices = environment.get("NVIDIA_VISIBLE_DEVICES", "").strip().lower()
|
||||
if visible_devices not in {"", "none", "void"}:
|
||||
return True
|
||||
return any((device_root / name).exists() for name in ("nvidiactl", "nvidia-uvm", "nvidia0"))
|
||||
|
||||
|
||||
def nvidia_required(
|
||||
mode: str,
|
||||
*,
|
||||
environ: dict[str, str] | None = None,
|
||||
device_root: Path = Path("/dev"),
|
||||
) -> bool:
|
||||
if mode == "nvidia":
|
||||
return True
|
||||
if mode == "cpu":
|
||||
return False
|
||||
return nvidia_runtime_observed(environ=environ, device_root=device_root)
|
||||
|
||||
|
||||
def _runtime_description() -> str:
|
||||
libc_name, libc_version = platform.libc_ver()
|
||||
if libc_name:
|
||||
return f"{libc_name} {libc_version}".strip()
|
||||
if Path("/lib/ld-musl-x86_64.so.1").exists():
|
||||
return "musl"
|
||||
return "unknown libc"
|
||||
|
||||
|
||||
def validate_nvidia_collection(
|
||||
collection: NvidiaCollection,
|
||||
*,
|
||||
required: bool,
|
||||
startup: bool = True,
|
||||
) -> None:
|
||||
"""Enforce the GPU-node startup/publication contract without fabricating observations."""
|
||||
|
||||
if not required:
|
||||
return
|
||||
setting = "MODELFORGE_AGENT_ACCELERATOR_MODE=nvidia"
|
||||
if collection.availability is Availability.TEMPORARILY_FAILED:
|
||||
reason = collection.reason or "NVML telemetry collection temporarily failed"
|
||||
problem = AgentStartupProblem(
|
||||
code=AgentStartupFailureCode.NVIDIA_TELEMETRY_UNAVAILABLE,
|
||||
setting=setting,
|
||||
message=f"NVML returned a temporary incomplete observation ({reason})",
|
||||
)
|
||||
if startup:
|
||||
raise AgentStartupError(problem)
|
||||
raise AgentObservationError(problem)
|
||||
if collection.availability is not Availability.KNOWN:
|
||||
reason = collection.reason or "NVML did not provide a reason"
|
||||
raise AgentStartupError(
|
||||
AgentStartupProblem(
|
||||
code=AgentStartupFailureCode.NVIDIA_NVML_UNAVAILABLE,
|
||||
setting=setting,
|
||||
message=(
|
||||
f"NVML is unavailable ({reason}) under {_runtime_description()}; "
|
||||
"verify the NVIDIA Container Toolkit, injected driver libraries, and "
|
||||
"/dev/nvidia* device mapping"
|
||||
),
|
||||
)
|
||||
)
|
||||
if not collection.inventory:
|
||||
raise AgentStartupError(
|
||||
AgentStartupProblem(
|
||||
code=AgentStartupFailureCode.NVIDIA_DEVICE_NOT_FOUND,
|
||||
setting=setting,
|
||||
message="NVML initialized but returned no NVIDIA accelerator inventory",
|
||||
)
|
||||
)
|
||||
inventory_uuids = {item.device_uuid for item in collection.inventory}
|
||||
telemetry_uuids = {item.device_uuid for item in collection.telemetry}
|
||||
missing = sorted(inventory_uuids - telemetry_uuids)
|
||||
if missing:
|
||||
raise AgentStartupError(
|
||||
AgentStartupProblem(
|
||||
code=AgentStartupFailureCode.NVIDIA_TELEMETRY_UNAVAILABLE,
|
||||
setting=setting,
|
||||
message="NVML returned no telemetry for accelerator(s): " + ", ".join(missing),
|
||||
)
|
||||
)
|
||||
@@ -0,0 +1,33 @@
|
||||
from pathlib import Path
|
||||
from typing import Literal
|
||||
|
||||
from pydantic import Field, SecretStr
|
||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
|
||||
|
||||
class AgentSettings(BaseSettings):
|
||||
model_config = SettingsConfigDict(
|
||||
env_prefix="MODELFORGE_AGENT_", env_file=".env", extra="ignore"
|
||||
)
|
||||
|
||||
control_plane_url: str = "http://api:8000"
|
||||
enrollment_token: SecretStr | None = None
|
||||
identity: str | None = None
|
||||
identity_mode: str = "auto"
|
||||
accelerator_mode: Literal["auto", "nvidia", "cpu"] = "auto"
|
||||
identity_file: Path = Path("/data/state/node-id")
|
||||
credential_file: Path = Path("/data/state/node-credential")
|
||||
state_file: Path = Path("/data/state/agent-state.json")
|
||||
display_name: str | None = None
|
||||
heartbeat_interval_seconds: int = Field(default=10, ge=1)
|
||||
inventory_interval_seconds: int = Field(default=300, ge=5)
|
||||
telemetry_interval_seconds: int = Field(default=30, ge=5)
|
||||
request_timeout_seconds: float = Field(default=10, gt=0)
|
||||
max_backoff_seconds: int = Field(default=60, ge=1)
|
||||
tls_verify: bool = True
|
||||
model_cache_path: Path = Path("/data/hf-cache")
|
||||
artifact_path: Path = Path("/data/artifacts")
|
||||
quarantine_path: Path = Path("/data/quarantine")
|
||||
hf_token: SecretStr | None = None
|
||||
artifact_job_poll_interval_seconds: int = Field(default=5, ge=1, le=60)
|
||||
artifact_download_timeout_seconds: float = Field(default=120, gt=0, le=600)
|
||||
@@ -0,0 +1,46 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from dataclasses import asdict, dataclass
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
@dataclass
|
||||
class AgentPersistentState:
|
||||
inventory_sequence: int = 0
|
||||
telemetry_sequence: int = 0
|
||||
|
||||
|
||||
class AgentStateStore:
|
||||
def __init__(self, state_file: Path, credential_file: Path) -> None:
|
||||
self.state_file = state_file
|
||||
self.credential_file = credential_file
|
||||
|
||||
def load(self) -> AgentPersistentState:
|
||||
try:
|
||||
payload = json.loads(self.state_file.read_text(encoding="utf-8"))
|
||||
return AgentPersistentState(
|
||||
inventory_sequence=int(payload.get("inventory_sequence", 0)),
|
||||
telemetry_sequence=int(payload.get("telemetry_sequence", 0)),
|
||||
)
|
||||
except (OSError, ValueError, TypeError, json.JSONDecodeError):
|
||||
return AgentPersistentState()
|
||||
|
||||
def save(self, state: AgentPersistentState) -> None:
|
||||
self.state_file.parent.mkdir(parents=True, exist_ok=True)
|
||||
temporary = self.state_file.with_suffix(".tmp")
|
||||
temporary.write_text(json.dumps(asdict(state), sort_keys=True), encoding="utf-8")
|
||||
os.replace(temporary, self.state_file)
|
||||
|
||||
def credential(self) -> str | None:
|
||||
try:
|
||||
value = self.credential_file.read_text(encoding="utf-8").strip()
|
||||
return value or None
|
||||
except OSError:
|
||||
return None
|
||||
|
||||
def save_credential(self, credential: str) -> None:
|
||||
self.credential_file.parent.mkdir(parents=True, exist_ok=True)
|
||||
self.credential_file.write_text(credential, encoding="utf-8")
|
||||
os.chmod(self.credential_file, 0o600)
|
||||
@@ -0,0 +1,63 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
|
||||
class AgentTransport:
|
||||
def __init__(self, base_url: str, timeout: float, verify: bool) -> None:
|
||||
self.client = httpx.Client(base_url=base_url.rstrip("/"), timeout=timeout, verify=verify)
|
||||
|
||||
def _request(
|
||||
self,
|
||||
method: str,
|
||||
path: str,
|
||||
payload: dict[str, Any] | None = None,
|
||||
credential: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
headers = {"Authorization": f"Bearer {credential}"} if credential else {}
|
||||
response = self.client.request(method, path, json=payload, headers=headers)
|
||||
response.raise_for_status()
|
||||
decoded = response.json()
|
||||
return dict(decoded) if isinstance(decoded, dict) else {}
|
||||
|
||||
def enroll(self, payload: dict[str, Any]) -> dict[str, Any]:
|
||||
return self._request("POST", "/api/v1/agent/enroll", payload)
|
||||
|
||||
def heartbeat(self, payload: dict[str, Any], credential: str) -> dict[str, Any]:
|
||||
return self._request("POST", "/api/v1/agent/heartbeat", payload, credential)
|
||||
|
||||
def inventory(self, payload: dict[str, Any], credential: str) -> dict[str, Any]:
|
||||
return self._request("PUT", "/api/v1/agent/inventory", payload, credential)
|
||||
|
||||
def telemetry(self, payload: dict[str, Any], credential: str) -> dict[str, Any]:
|
||||
return self._request("PUT", "/api/v1/agent/telemetry", payload, credential)
|
||||
|
||||
def next_artifact_job(self, credential: str) -> dict[str, Any] | None:
|
||||
response = self._request("GET", "/api/v1/agent/artifact-jobs/next", credential=credential)
|
||||
return response or None
|
||||
|
||||
def artifact_progress(
|
||||
self, job_id: str, payload: dict[str, Any], credential: str
|
||||
) -> dict[str, Any]:
|
||||
return self._request(
|
||||
"POST", f"/api/v1/agent/artifact-jobs/{job_id}/progress", payload, credential
|
||||
)
|
||||
|
||||
def artifact_complete(
|
||||
self, job_id: str, payload: dict[str, Any], credential: str
|
||||
) -> dict[str, Any]:
|
||||
return self._request(
|
||||
"POST", f"/api/v1/agent/artifact-jobs/{job_id}/complete", payload, credential
|
||||
)
|
||||
|
||||
def artifact_fail(
|
||||
self, job_id: str, payload: dict[str, Any], credential: str
|
||||
) -> dict[str, Any]:
|
||||
return self._request(
|
||||
"POST", f"/api/v1/agent/artifact-jobs/{job_id}/fail", payload, credential
|
||||
)
|
||||
|
||||
def close(self) -> None:
|
||||
self.client.close()
|
||||
Reference in New Issue
Block a user