202 lines
8.7 KiB
Python
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"]
|