import uuid from datetime import UTC, datetime import pytest from fastapi import HTTPException from fastapi.testclient import TestClient from hardware_fakes import FakeAcceleratorCollector, FakeHostCollector from pydantic import SecretStr from sqlalchemy import create_engine, text from sqlalchemy.orm import Session from sqlalchemy.pool import StaticPool from modelforge_api.api.routes.agent import AttemptLimiter, get_agent_service from modelforge_api.api.routes.hardware import get_hardware_service from modelforge_api.main import app from modelforge_api.persistence.models import Base, ComputeNode, NodeCredential, NodeEnrollment from modelforge_api.services.hardware_inventory import HardwareInventoryService from modelforge_api.services.node_agent import NodeAgentService from modelforge_api.settings import Settings, get_settings def enrollment_payload(token: str, identity: str = "remote-a") -> dict: return { "enrollment_token": token, "identity_key": identity, "identity_source": "persisted_uuid", "hostname": identity, "display_name": identity, "metadata": { "agent_version": "0.1.0", "protocol_version": 1, "supported_capabilities": ["hardware.inventory", "hardware.telemetry"], "started_at": datetime.now(UTC).isoformat(), }, } def test_agent_api_enrollment_secret_is_once_only_and_node_scoped() -> None: engine = create_engine( "sqlite+pysqlite:///:memory:", connect_args={"check_same_thread": False}, poolclass=StaticPool, ) Base.metadata.create_all(engine) session = Session(engine) settings = Settings(_env_file=None, operator_api_key=SecretStr("admin-key")) agent_service = NodeAgentService(session, settings) hardware_service = HardwareInventoryService( session, FakeHostCollector(), FakeAcceleratorCollector([]) ) app.dependency_overrides[get_agent_service] = lambda: agent_service app.dependency_overrides[get_hardware_service] = lambda: hardware_service app.dependency_overrides[get_settings] = lambda: settings try: with TestClient(app) as client: assert client.post("/api/v1/admin/node-enrollments", json={}).status_code == 401 created = client.post( "/api/v1/admin/node-enrollments", headers={"X-ModelForge-Admin-Token": "admin-key"}, json={"display_name": "Remote A", "role": "primary-inference"}, ) assert created.status_code == 201 assert created.headers["cache-control"] == "no-store" token = created.json()["enrollment_token"] summaries = client.get( "/api/v1/admin/node-enrollments", headers={"X-ModelForge-Admin-Token": "admin-key"}, ).json() assert "enrollment_token" not in summaries[0] enrolled = client.post("/api/v1/agent/enroll", json=enrollment_payload(token)) assert enrolled.status_code == 201 assert enrolled.headers["cache-control"] == "no-store" credential = enrolled.json()["node_credential"] assert ( client.post("/api/v1/agent/enroll", json=enrollment_payload(token)).status_code == 401 ) wrong_node = client.post( "/api/v1/agent/heartbeat", headers={"Authorization": f"Bearer {credential}"}, json={ "identity_key": "remote-b", "metadata": enrollment_payload(token)["metadata"], "observed_at": datetime.now(UTC).isoformat(), }, ) assert wrong_node.status_code == 403 assert wrong_node.json()["error"]["code"] == "agent_not_authorized" accepted = client.post( "/api/v1/agent/heartbeat", headers={"Authorization": f"Bearer {credential}"}, json={ "identity_key": "remote-a", "metadata": enrollment_payload(token)["metadata"], "observed_at": datetime.now(UTC).isoformat(), }, ) assert accepted.status_code == 200 assert accepted.json()["accepted"] is True host_inventory = ( FakeHostCollector() .collect() .model_copy( update={ "identity_key": "remote-a", "identity_source": "persisted_uuid", "hostname": "remote-a", "display_name": "remote-a", } ) ) published_inventory = client.put( "/api/v1/agent/inventory", headers={"Authorization": f"Bearer {credential}"}, json={ "identity_key": "remote-a", "protocol_version": 1, "sequence": 1, "observed_at": host_inventory.inventory_at.isoformat(), "host": host_inventory.model_dump(mode="json"), "nvidia": {"availability": "known", "inventory": []}, }, ) assert published_inventory.status_code == 200 assert published_inventory.json()["accepted"] is True wrong_scope_created = client.post( "/api/v1/admin/node-enrollments", headers={"X-ModelForge-Admin-Token": "admin-key"}, json={"display_name": "Wrong scope control"}, ) assert wrong_scope_created.status_code == 201 wrong_scope_enrollment = session.get( NodeEnrollment, uuid.UUID(wrong_scope_created.json()["id"]), ) assert wrong_scope_enrollment is not None session.execute(text("PRAGMA ignore_check_constraints = ON")) wrong_scope_enrollment.scope = "node.publish" session.commit() session.execute(text("PRAGMA ignore_check_constraints = OFF")) rejected_enrollment = client.post( "/api/v1/agent/enroll", json=enrollment_payload( wrong_scope_created.json()["enrollment_token"], identity="wrong-scope-node", ), ) assert rejected_enrollment.status_code == 403 assert rejected_enrollment.json()["error"]["code"] == "agent_not_authorized" session.refresh(wrong_scope_enrollment) assert wrong_scope_enrollment.used_at is None assert session.query(ComputeNode).count() == 1 assert session.query(NodeCredential).count() == 1 node_id = enrolled.json()["node_id"] node = client.get( f"/api/v1/hardware/nodes/{node_id}", headers={"X-ModelForge-Admin-Token": "admin-key"}, ) assert node.status_code == 200 assert node.json()["role"] == "primary-inference" assert node.json()["observation_source"] == "remote_agent" malformed = client.put( "/api/v1/agent/telemetry", headers={"Authorization": f"Bearer {credential}"}, json={"identity_key": "remote-a"}, ) assert malformed.status_code == 422 assert malformed.json()["error"]["code"] == "request_validation_failed" finally: app.dependency_overrides.clear() session.close() def test_enrollment_attempt_limiter_rejects_bursts() -> None: limiter = AttemptLimiter(limit=2, window_seconds=60) limiter.check("test-client") limiter.check("test-client") with pytest.raises(HTTPException) as raised: limiter.check("test-client") assert raised.value.status_code == 429 def test_agent_openapi_contract_contains_versioned_surfaces_without_list_secret() -> None: schema = app.openapi() for path in ( "/api/v1/agent/enroll", "/api/v1/agent/heartbeat", "/api/v1/agent/inventory", "/api/v1/agent/telemetry", "/api/v1/admin/node-enrollments", "/api/v1/admin/hardware/nodes/{node_id}/credential/rotate", "/api/v1/admin/hardware/nodes/{node_id}/decommission/preview", "/api/v1/admin/hardware/nodes/{node_id}/decommission", "/api/v1/agent/runtime-probes/next", "/api/v1/agent/runtime-probes/{probe_id}/progress", "/api/v1/agent/runtime-probes/{probe_id}/complete", "/api/v1/agent/runtime-probes/{probe_id}/fail", ): assert path in schema["paths"] summary = schema["components"]["schemas"]["EnrollmentTokenSummary"] assert "enrollment_token" not in summary["properties"]