Files
ModelForge/backend/tests/test_node_agent_api.py

202 lines
8.7 KiB
Python

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"]