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()