Files
ModelForge/backend/tests/test_node_decommission.py

538 lines
19 KiB
Python

from __future__ import annotations
import uuid
from datetime import UTC, datetime
import pytest
from fastapi import FastAPI
from fastapi.testclient import TestClient
from pydantic import SecretStr
from sqlalchemy import create_engine, func, select
from sqlalchemy.orm import Session
from sqlalchemy.pool import StaticPool
from modelforge_api.api.routes.agent import get_decommission_service
from modelforge_api.api.routes.agent import router as agent_router
from modelforge_api.domain.agent_protocol import (
AGENT_PROTOCOL_CAPABILITIES,
AgentMetadata,
EnrollmentRequest,
EnrollmentTokenCreate,
HeartbeatRequest,
NodeManagementUpdate,
)
from modelforge_api.domain.node_decommission import NodeDecommissionExecute
from modelforge_api.main import node_decommission_error
from modelforge_api.persistence.models import (
Accelerator,
AcceleratorTelemetryLatest,
ArtifactJob,
AuditEvent,
Base,
CapabilityDeployment,
ComputeNode,
GatewayRequest,
GpuLease,
HardwareInventoryRun,
HostTelemetryLatest,
LifecycleApprovalRequest,
NodeCredential,
NodeDecommissionOperation,
ResidencyAllocation,
RuntimeProbe,
SchedulerAcceleratorState,
ServingJob,
StorageVolumeState,
)
from modelforge_api.services.invariants import check_invariants
from modelforge_api.services.manifest_registry import ManifestRegistry
from modelforge_api.services.node_agent import (
AgentAuthenticationError,
AgentConflictError,
NodeAgentService,
secret_hash,
)
from modelforge_api.services.node_decommission import (
NodeDecommissionError,
NodeDecommissionService,
)
from modelforge_api.services.serving import ServingService
from modelforge_api.services.transient_payloads import MemoryPayloadStore
from modelforge_api.settings import Settings, get_settings
@pytest.fixture
def session() -> Session:
engine = create_engine("sqlite+pysqlite:///:memory:")
Base.metadata.create_all(engine)
with Session(engine) as value:
yield value
def _node(session: Session, *, liveness: str = "offline") -> ComputeNode:
node = ComputeNode(
key=f"disposable-{uuid.uuid4()}",
hostname="disposable-node",
display_name="Disposable node",
enabled=True,
status="active",
liveness_state=liveness,
inventory={"source": "test"},
agent_capabilities=["hardware.inventory"],
)
session.add(node)
session.commit()
return node
def _request(preview, *, confirmation: str | None = None) -> NodeDecommissionExecute:
return NodeDecommissionExecute(
expected_generation=preview.node_generation,
preview_digest=preview.dependency_digest,
idempotency_key=f"decom-{uuid.uuid4()}",
operator="test-operator",
reason="Disposable acceptance node has permanently left service.",
confirmation=confirmation or preview.persisted_identity,
)
def _blocker_codes(service: NodeDecommissionService, node: ComputeNode) -> set[str]:
return {item.code for item in service.preview(node.id).blockers}
def test_safe_offline_unused_node_decommissions_once_and_retains_provenance(
session: Session,
) -> None:
node = _node(session)
credential = NodeCredential(
compute_node_id=node.id,
secret_hash="a" * 64,
scope="node.publish",
)
history = HardwareInventoryRun(
compute_node_id=node.id,
status="completed",
source="remote_agent",
fingerprint="b" * 64,
summary={"retained": True},
started_at=datetime.now(UTC),
completed_at=datetime.now(UTC),
)
session.add_all([credential, history])
session.commit()
service = NodeDecommissionService(session)
preview = service.preview(node.id)
assert preview.safe is True
result = service.execute(node.id, _request(preview))
session.refresh(node)
session.refresh(credential)
assert result.status == "completed"
assert result.idempotent_replay is False
assert result.credential_revocations == 1
assert credential.revoked_at is not None
assert node.decommissioned_at is not None
assert node.enabled is False
assert node.status == "decommissioned"
assert node.liveness_state == "decommissioned"
assert node.inventory == {}
assert node.agent_capabilities == []
assert session.get(HardwareInventoryRun, history.id) is not None
assert session.get(ComputeNode, node.id) is node
replay = service.execute(node.id, _request(preview))
assert replay.operation_id == result.operation_id
assert replay.idempotent_replay is True
assert session.scalar(select(func.count()).select_from(NodeDecommissionOperation)) == 1
assert (
session.scalar(
select(func.count())
.select_from(AuditEvent)
.where(AuditEvent.action == "NODE_DECOMMISSIONED")
)
== 1
)
def test_online_node_is_refused_without_mutation(session: Session) -> None:
node = _node(session, liveness="online")
service = NodeDecommissionService(session)
preview = service.preview(node.id)
assert preview.safe is False
assert "node_online" in {item.code for item in preview.blockers}
with pytest.raises(NodeDecommissionError, match="preconditions") as raised:
service.execute(node.id, _request(preview))
assert raised.value.code == "node_decommission_blocked"
session.refresh(node)
assert node.decommissioned_at is None
assert node.enabled is True
@pytest.mark.parametrize(
("record_factory", "expected_code"),
[
(
lambda node: ArtifactJob(
plan_id=uuid.uuid4(),
compute_node_id=node.id,
storage_root_id=uuid.uuid4(),
status="queued",
idempotency_key=str(uuid.uuid4()),
total_bytes=1,
),
"active_artifact_job",
),
(
lambda node: ServingJob(
capability_deployment_id=uuid.uuid4(),
compute_node_id=node.id,
operation="generate",
status="running",
idempotency_key=str(uuid.uuid4()),
),
"active_serving_job",
),
(
lambda node: GatewayRequest(
request_id=uuid.uuid4(),
capability_key="assistant.general",
capability_version=1,
compute_node_id=node.id,
status="queued",
priority="production",
input_sha256="c" * 64,
input_count=1,
),
"active_gateway_request",
),
(
lambda node: RuntimeProbe(
artifact_set_id=uuid.uuid4(),
runtime_profile_id=uuid.uuid4(),
compute_node_id=node.id,
compatibility_assessment_id=uuid.uuid4(),
execution_approval_id=uuid.uuid4(),
status="loading",
probe_input="test",
idempotency_key=str(uuid.uuid4()),
environment_fingerprint="d" * 64,
),
"active_runtime_probe",
),
(
lambda node: LifecycleApprovalRequest(
policy_revision_id=uuid.uuid4(),
target_type="compute_node",
target_ref=str(node.id),
environment="production",
requested_transition="promote",
evidence_snapshot={"compute_node_id": str(node.id)},
evidence_fingerprint="e" * 64,
status="PENDING",
requested_by="operator",
reason="test dependency",
),
"pending_lifecycle_approval",
),
],
ids=["artifact-job", "serving-job", "gateway-request", "runtime-probe", "lifecycle"],
)
def test_active_work_classes_block_decommission(
session: Session, record_factory, expected_code: str
) -> None:
node = _node(session)
session.add(record_factory(node))
session.commit()
assert expected_code in _blocker_codes(NodeDecommissionService(session), node)
def test_gpu_lease_and_runtime_residency_each_block(session: Session) -> None:
node = _node(session)
accelerator = Accelerator(
compute_node_id=node.id,
device_index=0,
device_uuid="GPU-DISPOSABLE",
name="Disposable GPU",
status="active",
)
session.add(accelerator)
session.commit()
lease = GpuLease(
accelerator_id=accelerator.id,
deployment_id=uuid.uuid4(),
priority="production",
reserved_vram_mb=1024,
state="active",
)
residency = ResidencyAllocation(
capability_deployment_id=uuid.uuid4(),
compute_node_id=node.id,
accelerator_id=accelerator.id,
state="ready",
)
session.add_all([lease, residency])
session.commit()
codes = _blocker_codes(NodeDecommissionService(session), node)
assert {"active_gpu_lease", "active_runtime_residency"} <= codes
def test_active_production_deployment_blocks_decommission(session: Session) -> None:
node = _node(session)
session.add(
CapabilityDeployment(
capability_contract_id=uuid.uuid4(),
deployment_candidate_id=uuid.uuid4(),
artifact_set_id=uuid.uuid4(),
runtime_profile_id=uuid.uuid4(),
compute_node_id=node.id,
accelerator_id=uuid.uuid4(),
status="stable",
production=True,
config_fingerprint="f" * 64,
provenance={"test": "disposable"},
)
)
session.commit()
assert "active_production_deployment" in _blocker_codes(NodeDecommissionService(session), node)
def test_new_work_after_preview_forces_fresh_preview(session: Session) -> None:
node = _node(session)
service = NodeDecommissionService(session)
preview = service.preview(node.id)
session.add(
ServingJob(
capability_deployment_id=uuid.uuid4(),
compute_node_id=node.id,
operation="generate",
status="queued",
idempotency_key=str(uuid.uuid4()),
)
)
session.commit()
with pytest.raises(NodeDecommissionError, match="generate a new preview") as raised:
service.execute(node.id, _request(preview))
assert raised.value.code == "decommission_preview_stale"
assert raised.value.details["preview"]["safe"] is False
def test_confirmation_and_unknown_node_fail_closed(session: Session) -> None:
node = _node(session)
service = NodeDecommissionService(session)
preview = service.preview(node.id)
with pytest.raises(NodeDecommissionError) as mismatch:
service.execute(node.id, _request(preview, confirmation="wrong-node"))
assert mismatch.value.code == "decommission_confirmation_mismatch"
with pytest.raises(NodeDecommissionError) as missing:
service.preview(uuid.uuid4())
assert missing.value.status_code == 404
assert missing.value.code == "node_not_found"
def test_current_scheduler_and_telemetry_truth_is_removed_but_identity_remains(
session: Session,
) -> None:
node = _node(session)
accelerator = Accelerator(
compute_node_id=node.id,
device_index=0,
device_uuid="GPU-CLEANUP",
name="Cleanup GPU",
status="active",
)
session.add(accelerator)
session.flush()
session.add_all(
[
HostTelemetryLatest(
compute_node_id=node.id,
available_ram_bytes=1,
observed_at=datetime.now(UTC),
),
StorageVolumeState(
compute_node_id=node.id,
purpose="models",
path="/disposable",
total_bytes=10,
used_bytes=1,
free_bytes=9,
observed_at=datetime.now(UTC),
),
AcceleratorTelemetryLatest(
accelerator_id=accelerator.id,
observed_at=datetime.now(UTC),
),
SchedulerAcceleratorState(
accelerator_id=accelerator.id,
pressure_state="NORMAL",
last_observed_at=datetime.now(UTC),
),
]
)
session.commit()
service = NodeDecommissionService(session)
preview = service.preview(node.id)
assert preview.safe is True
service.execute(node.id, _request(preview))
assert session.scalar(select(func.count()).select_from(HostTelemetryLatest)) == 0
assert session.scalar(select(func.count()).select_from(StorageVolumeState)) == 0
assert session.scalar(select(func.count()).select_from(AcceleratorTelemetryLatest)) == 0
assert session.scalar(select(func.count()).select_from(SchedulerAcceleratorState)) == 0
session.refresh(accelerator)
assert accelerator.status == "decommissioned"
assert session.get(ComputeNode, node.id) is not None
def test_terminal_node_rejects_management_and_has_holding_invariant(session: Session) -> None:
node = _node(session)
service = NodeDecommissionService(session)
preview = service.preview(node.id)
service.execute(node.id, _request(preview))
settings = Settings(_env_file=None, operator_api_key=SecretStr("admin-key"))
agent = NodeAgentService(session, settings)
with pytest.raises(AgentConflictError, match="decommissioned"):
agent.update_node(node.id, NodeManagementUpdate(enabled=True))
with pytest.raises(AgentConflictError, match="decommissioned"):
agent.rotate_credential(node.id)
invariant = next(
item
for item in check_invariants(session).results
if item.key == "decommissioned_nodes_are_terminal"
)
assert invariant.status.value == "HOLDS"
def test_old_credential_is_unusable_after_decommission(session: Session) -> None:
node = _node(session)
credential_id = uuid.uuid4()
raw = f"mfnode_{credential_id}_disposable-secret"
session.add(
NodeCredential(
id=credential_id,
compute_node_id=node.id,
secret_hash=secret_hash(raw),
scope="node.publish",
)
)
session.commit()
service = NodeDecommissionService(session)
preview = service.preview(node.id)
service.execute(node.id, _request(preview))
agent = NodeAgentService(session, Settings(_env_file=None))
with pytest.raises(AgentAuthenticationError, match="revoked or unknown"):
agent.authenticate(f"Bearer {raw}")
def test_old_agent_report_cannot_resurrect_tombstone(session: Session) -> None:
node = _node(session)
service = NodeDecommissionService(session)
preview = service.preview(node.id)
service.execute(node.id, _request(preview))
agent = NodeAgentService(session, Settings(_env_file=None))
metadata = AgentMetadata(
agent_version="1.0.0",
protocol_version=1,
supported_capabilities=AGENT_PROTOCOL_CAPABILITIES,
started_at=datetime.now(UTC),
)
with pytest.raises(AgentConflictError, match="decommissioned"):
agent.heartbeat(
node,
HeartbeatRequest(
identity_key=node.key,
metadata=metadata,
observed_at=datetime.now(UTC),
),
)
session.refresh(node)
assert node.liveness_state == "decommissioned"
assert node.enabled is False
def test_old_persisted_identity_cannot_ordinarily_reenroll(session: Session) -> None:
node = _node(session)
service = NodeDecommissionService(session)
preview = service.preview(node.id)
service.execute(node.id, _request(preview))
agent = NodeAgentService(session, Settings(_env_file=None))
token = agent.create_enrollment(EnrollmentTokenCreate(display_name="Disposable replacement"))
with pytest.raises(AgentConflictError, match="explicit recovery enrollment"):
agent.enroll(
EnrollmentRequest(
enrollment_token=token.enrollment_token,
identity_key=node.key,
identity_source="persisted_uuid",
hostname=node.hostname,
display_name=node.display_name,
metadata=AgentMetadata(
agent_version="1.0.0",
protocol_version=1,
supported_capabilities=AGENT_PROTOCOL_CAPABILITIES,
started_at=datetime.now(UTC),
),
)
)
def test_scheduler_placement_evidence_excludes_decommissioned_node(session: Session) -> None:
node = _node(session)
service = NodeDecommissionService(session)
preview = service.preview(node.id)
service.execute(node.id, _request(preview))
deployment = CapabilityDeployment(
capability_contract_id=uuid.uuid4(),
deployment_candidate_id=uuid.uuid4(),
artifact_set_id=uuid.uuid4(),
runtime_profile_id=uuid.uuid4(),
compute_node_id=node.id,
accelerator_id=uuid.uuid4(),
config_fingerprint="9" * 64,
provenance={},
)
evidence = ServingService(
session,
Settings(_env_file=None),
ManifestRegistry(),
MemoryPayloadStore(),
)._placement_evidence(deployment)
assert str(node.id) not in {candidate["node_id"] for candidate in evidence["candidate_nodes"]}
def test_admin_api_requires_operator_token_and_never_accepts_bearer_project_token() -> None:
engine = create_engine(
"sqlite+pysqlite:///:memory:",
connect_args={"check_same_thread": False},
poolclass=StaticPool,
)
Base.metadata.create_all(engine)
session = Session(engine)
node = _node(session)
settings = Settings(_env_file=None, operator_api_key=SecretStr("admin-key"))
service = NodeDecommissionService(session)
test_app = FastAPI()
test_app.include_router(agent_router)
test_app.add_exception_handler(NodeDecommissionError, node_decommission_error)
test_app.dependency_overrides[get_decommission_service] = lambda: service
test_app.dependency_overrides[get_settings] = lambda: settings
try:
with TestClient(test_app) as client:
path = f"/api/v1/admin/hardware/nodes/{node.id}/decommission/preview"
assert client.post(path).status_code == 401
assert (
client.post(path, headers={"Authorization": "Bearer project-token"}).status_code
== 401
)
accepted = client.post(path, headers={"X-ModelForge-Admin-Token": "admin-key"})
assert accepted.status_code == 200
assert accepted.json()["safe"] is True
missing = client.post(
f"/api/v1/admin/hardware/nodes/{uuid.uuid4()}/decommission/preview",
headers={"X-ModelForge-Admin-Token": "admin-key"},
)
assert missing.status_code == 404
assert missing.json()["error"]["code"] == "node_not_found"
finally:
test_app.dependency_overrides.clear()
session.close()