Files
ModelForge/backend/tests/test_node_agent_control_plane.py

600 lines
23 KiB
Python

import uuid
from datetime import UTC, datetime, timedelta
import pytest
from hardware_fakes import FakeHostCollector, accelerator
from pydantic import SecretStr
from sqlalchemy import create_engine, select, text, update
from sqlalchemy.orm import Session
from modelforge_api.domain.agent_protocol import (
AGENT_PROTOCOL_CAPABILITIES,
AgentMetadata,
EnrollmentRequest,
EnrollmentTokenCreate,
HeartbeatRequest,
InventoryNvidiaPayload,
InventoryReport,
TelemetryReport,
)
from modelforge_api.domain.enums import Availability, HardwareStatus, NodeLiveness
from modelforge_api.domain.hardware import AcceleratorTelemetry, ObservedValue
from modelforge_api.persistence.models import (
Accelerator,
AuditEvent,
Base,
ComputeNode,
NodeCredential,
NodeEnrollment,
)
from modelforge_api.services.node_agent import (
AgentAuthenticationError,
AgentAuthorizationError,
AgentConflictError,
NodeAgentService,
secret_hash,
)
from modelforge_api.settings import Settings
def settings() -> Settings:
return Settings(
_env_file=None,
operator_api_key=SecretStr("admin-test-key"),
node_stale_after_seconds=10,
node_offline_after_seconds=20,
)
def metadata(protocol: int = 1) -> AgentMetadata:
return AgentMetadata(
agent_version="0.1.0",
protocol_version=protocol,
supported_capabilities=AGENT_PROTOCOL_CAPABILITIES,
started_at=datetime.now(UTC),
)
def enroll(service: NodeAgentService, identity: str = "remote-node"):
created = service.create_enrollment(
EnrollmentTokenCreate(
display_name="Server 01",
role="primary-inference",
labels={"site": "lab"},
production_eligible=True,
lab_eligible=True,
benchmark_eligible=True,
)
)
response = service.enroll(
EnrollmentRequest(
enrollment_token=created.enrollment_token,
identity_key=identity,
identity_source="persisted_uuid",
hostname="server01",
display_name="server01",
metadata=metadata(),
)
)
return created, response
def inventory(sequence: int, devices=None, hostname: str = "server01") -> InventoryReport:
host = (
FakeHostCollector()
.collect()
.model_copy(
update={
"identity_key": "remote-node",
"identity_source": "persisted_uuid",
"hostname": hostname,
"display_name": hostname,
}
)
)
return InventoryReport(
identity_key="remote-node",
protocol_version=1,
sequence=sequence,
observed_at=host.inventory_at + timedelta(seconds=sequence),
host=host,
nvidia=InventoryNvidiaPayload(
availability=Availability.KNOWN,
inventory=devices if devices is not None else [accelerator()],
),
)
def telemetry(observed_at: datetime, sequence: int) -> TelemetryReport:
host = FakeHostCollector().collect()
known_int = ObservedValue.known
return TelemetryReport(
identity_key="remote-node",
protocol_version=1,
sequence=sequence,
observed_at=observed_at,
available_ram_bytes=host.available_ram_bytes,
storage=host.storage,
accelerators=[
AcceleratorTelemetry(
device_uuid="GPU-A",
observed_at=observed_at,
used_vram_bytes=known_int(1024),
free_vram_bytes=known_int(2048),
gpu_utilization_percent=known_int(42),
memory_utilization_percent=known_int(12),
temperature_c=known_int(55),
power_draw_w=ObservedValue.known(80.0),
power_limit_w=ObservedValue.known(200.0),
graphics_clock_mhz=known_int(1500),
memory_clock_mhz=known_int(7000),
fan_speed_percent=ObservedValue.absent(Availability.UNSUPPORTED, "not supported"),
performance_state=ObservedValue.known("P2"),
)
],
)
@pytest.fixture
def session() -> Session:
engine = create_engine("sqlite+pysqlite:///:memory:")
Base.metadata.create_all(engine)
with Session(engine) as value:
yield value
def test_one_time_enrollment_hashes_secrets_and_audits(session: Session) -> None:
service = NodeAgentService(session, settings())
created, response = enroll(service)
record = session.scalar(select(NodeEnrollment))
credential = session.scalar(select(NodeCredential))
assert record is not None and created.enrollment_token not in record.token_hash
assert credential is not None and response.node_credential not in credential.secret_hash
assert len(record.token_hash) == 64 and len(credential.secret_hash) == 64
with pytest.raises(AgentAuthenticationError, match="already used"):
service.enroll(
EnrollmentRequest(
enrollment_token=created.enrollment_token,
identity_key="remote-node",
identity_source="persisted_uuid",
hostname="server01",
display_name="server01",
metadata=metadata(),
)
)
actions = list(session.scalars(select(AuditEvent.action)))
assert "NODE_ENROLLMENT_TOKEN_CREATED" in actions
assert "NODE_ENROLLED" in actions
second_token = service.create_enrollment(EnrollmentTokenCreate())
second = service.enroll(
EnrollmentRequest(
enrollment_token=second_token.enrollment_token,
identity_key="remote-node",
identity_source="persisted_uuid",
hostname="renamed-server",
display_name="renamed-server",
metadata=metadata(),
)
)
assert second.node_id == response.node_id
assert second.node_credential != response.node_credential
assert session.query(ComputeNode).count() == 1
with pytest.raises(AgentAuthenticationError, match="revoked"):
service.authenticate(f"Bearer {response.node_credential}")
def test_invalid_expired_and_revoked_enrollment_tokens_are_rejected(session: Session) -> None:
service = NodeAgentService(session, settings())
request = EnrollmentRequest(
enrollment_token="x" * 40,
identity_key="remote-node",
identity_source="persisted_uuid",
hostname="server01",
display_name="server01",
metadata=metadata(),
)
with pytest.raises(AgentAuthenticationError, match="invalid"):
service.enroll(request)
created = service.create_enrollment(EnrollmentTokenCreate())
record = session.scalar(select(NodeEnrollment))
assert record is not None
record.expires_at = datetime.now(UTC) - timedelta(seconds=1)
session.commit()
with pytest.raises(AgentAuthenticationError, match="expired"):
service.enroll(request.model_copy(update={"enrollment_token": created.enrollment_token}))
incompatible = service.create_enrollment(EnrollmentTokenCreate())
with pytest.raises(AgentConflictError, match="incompatible"):
service.enroll(
request.model_copy(
update={
"enrollment_token": incompatible.enrollment_token,
"metadata": metadata().model_copy(update={"protocol_version": 99}),
}
)
)
revoked = service.create_enrollment(EnrollmentTokenCreate())
service.revoke_enrollment(uuid.UUID(revoked.id))
with pytest.raises(AgentAuthenticationError, match="revoked"):
service.enroll(request.model_copy(update={"enrollment_token": revoked.enrollment_token}))
def test_revocation_between_validation_and_claim_wins_without_minting(
session: Session, monkeypatch: pytest.MonkeyPatch
) -> None:
"""Trigger the reviewer's validation/revoke/claim interleaving deterministically."""
service = NodeAgentService(session, settings())
created = service.create_enrollment(EnrollmentTokenCreate())
enrollment_id = uuid.UUID(created.id)
original_claim = service._claim_enrollment
def revoke_then_claim(
*, enrollment_id: uuid.UUID, token_hash: str, claim_now: datetime
) -> bool:
service.revoke_enrollment(enrollment_id)
return original_claim(
enrollment_id=enrollment_id,
token_hash=token_hash,
claim_now=claim_now,
)
monkeypatch.setattr(service, "_claim_enrollment", revoke_then_claim)
with pytest.raises(AgentAuthenticationError, match="revoked"):
service.enroll(
EnrollmentRequest(
enrollment_token=created.enrollment_token,
identity_key="revocation-race-node",
identity_source="persisted_uuid",
hostname="revocation-race-node",
display_name="revocation-race-node",
metadata=metadata(),
)
)
enrollment = session.get(NodeEnrollment, enrollment_id, populate_existing=True)
assert enrollment is not None
assert enrollment.revoked_at is not None
assert enrollment.used_at is None
assert session.query(ComputeNode).count() == 0
assert session.query(NodeCredential).count() == 0
assert (
session.query(AuditEvent)
.filter(AuditEvent.action == "NODE_ENROLLMENT_TOKEN_REVOKED")
.count()
== 1
)
# The already-revoked DELETE control stays idempotent and must not duplicate its audit.
service.revoke_enrollment(enrollment_id)
assert (
session.query(AuditEvent)
.filter(AuditEvent.action == "NODE_ENROLLMENT_TOKEN_REVOKED")
.count()
== 1
)
def test_a_successful_claim_makes_a_later_revocation_lose_without_revoke_audit(
session: Session,
) -> None:
service = NodeAgentService(session, settings())
created, response = enroll(service, identity="claim-winner-node")
with pytest.raises(AgentConflictError, match="used"):
service.revoke_enrollment(uuid.UUID(created.id))
enrollment = session.get(NodeEnrollment, uuid.UUID(created.id), populate_existing=True)
credential = session.get(NodeCredential, uuid.UUID(response.credential_id))
assert enrollment is not None and enrollment.used_at is not None
assert enrollment.revoked_at is None
assert credential is not None and credential.revoked_at is None
assert (
session.query(AuditEvent)
.filter(AuditEvent.action == "NODE_ENROLLMENT_TOKEN_REVOKED")
.count()
== 0
)
@pytest.mark.parametrize(
("mutation", "error_type", "message"),
[
("expired", AgentAuthenticationError, "expired"),
("used", AgentAuthenticationError, "already used"),
("scope", AgentAuthorizationError, "scope"),
("token_identity", AgentAuthenticationError, "invalid"),
],
)
def test_every_mutable_authorization_fact_is_rechecked_at_the_atomic_claim(
session: Session,
monkeypatch: pytest.MonkeyPatch,
mutation: str,
error_type: type[Exception],
message: str,
) -> None:
"""Controlled SQLite interleavings cover expiry, use, scope, and token-identity TOCTOU."""
service = NodeAgentService(session, settings())
created = service.create_enrollment(EnrollmentTokenCreate())
enrollment_id = uuid.UUID(created.id)
original_claim = service._claim_enrollment
def mutate_then_claim(
*, enrollment_id: uuid.UUID, token_hash: str, claim_now: datetime
) -> bool:
values: dict[str, object]
if mutation == "expired":
values = {"expires_at": claim_now - timedelta(microseconds=1)}
elif mutation == "used":
values = {"used_at": claim_now}
elif mutation == "scope":
# Production schema 0023 rejects this write. Temporarily bypass SQLite's constraint
# to prove the claim predicate still fails closed for a malformed legacy/interleaved
# row; the migration suite separately proves the real constraint is enforced.
session.execute(text("PRAGMA ignore_check_constraints = ON"))
values = {"scope": "node.publish"}
else:
values = {"token_hash": secret_hash("a different enrollment token")}
session.execute(
update(NodeEnrollment).where(NodeEnrollment.id == enrollment_id).values(**values)
)
session.commit()
if mutation == "scope":
session.execute(text("PRAGMA ignore_check_constraints = OFF"))
return original_claim(
enrollment_id=enrollment_id,
token_hash=token_hash,
claim_now=claim_now,
)
monkeypatch.setattr(service, "_claim_enrollment", mutate_then_claim)
with pytest.raises(error_type, match=message):
service.enroll(
EnrollmentRequest(
enrollment_token=created.enrollment_token,
identity_key=f"{mutation}-race-node",
identity_source="persisted_uuid",
hostname=f"{mutation}-race-node",
display_name=f"{mutation}-race-node",
metadata=metadata(),
)
)
enrollment = session.get(NodeEnrollment, enrollment_id, populate_existing=True)
assert enrollment is not None
assert enrollment.enrolled_node_id is None
assert session.query(ComputeNode).count() == 0
assert session.query(NodeCredential).count() == 0
assert (
session.query(AuditEvent).filter(AuditEvent.action == "NODE_ENROLLED").count() == 0
)
@pytest.mark.parametrize(
("enrollment_scope", "publisher_scope"),
[
("node.publish", "node.enroll"),
("NODE.ENROLL", "NODE.PUBLISH"),
("node.enroll ", "node.publish "),
("node.enroll\u200b", "node.publish\u200b"),
],
)
def test_wrong_scope_enrollment_and_publisher_records_fail_closed_before_state_changes(
session: Session,
enrollment_scope: str,
publisher_scope: str,
) -> None:
service = NodeAgentService(session, settings())
created = service.create_enrollment(EnrollmentTokenCreate())
enrollment = session.get(NodeEnrollment, uuid.UUID(created.id))
assert enrollment is not None
session.execute(text("PRAGMA ignore_check_constraints = ON"))
enrollment.scope = enrollment_scope
session.commit()
session.execute(text("PRAGMA ignore_check_constraints = OFF"))
request = EnrollmentRequest(
enrollment_token=created.enrollment_token,
identity_key="wrong-scope-node",
identity_source="persisted_uuid",
hostname="wrong-scope-node",
display_name="wrong-scope-node",
metadata=metadata(),
)
with pytest.raises(AgentAuthorizationError, match="scope"):
service.enroll(request)
session.refresh(enrollment)
assert enrollment.used_at is None
assert session.query(ComputeNode).count() == 0
assert session.query(NodeCredential).count() == 0
enrollment.scope = "node.enroll"
session.commit()
response = service.enroll(request)
credential = session.get(NodeCredential, uuid.UUID(response.credential_id))
assert credential is not None
session.execute(text("PRAGMA ignore_check_constraints = ON"))
credential.scope = publisher_scope
session.commit()
session.execute(text("PRAGMA ignore_check_constraints = OFF"))
with pytest.raises(AgentAuthorizationError, match="scope"):
service.authenticate(f"Bearer {response.node_credential}")
def test_credential_ownership_disable_and_revocation(session: Session) -> None:
service = NodeAgentService(session, settings())
_created, response = enroll(service)
_credential, node = service.authenticate(f"Bearer {response.node_credential}")
heartbeat = HeartbeatRequest(
identity_key="another-node",
metadata=metadata(),
observed_at=datetime.now(UTC),
)
with pytest.raises(AgentAuthorizationError, match="another node"):
service.heartbeat(node, heartbeat)
node.enabled = False
session.commit()
with pytest.raises(AgentAuthorizationError, match="disabled"):
service.authenticate(f"Bearer {response.node_credential}")
node.enabled = True
session.commit()
rotated = service.rotate_credential(node.id)
assert rotated.node_credential != response.node_credential
with pytest.raises(AgentAuthenticationError, match="revoked"):
service.authenticate(f"Bearer {response.node_credential}")
service.authenticate(f"Bearer {rotated.node_credential}")
service.revoke_credential(node.id)
with pytest.raises(AgentAuthenticationError, match="revoked"):
service.authenticate(f"Bearer {rotated.node_credential}")
def test_remote_reconciliation_is_idempotent_non_destructive_and_ordered(
session: Session,
) -> None:
service = NodeAgentService(session, settings())
_created, response = enroll(service)
_credential, node = service.authenticate(f"Bearer {response.node_credential}")
assert service.publish_inventory(node, inventory(1)).accepted
first_accelerator = session.scalar(select(Accelerator))
assert first_accelerator is not None
first_id = first_accelerator.id
assert not service.publish_inventory(node, inventory(1)).accepted
assert service.publish_inventory(node, inventory(2, [], hostname="renamed-host")).accepted
session.refresh(first_accelerator)
assert first_accelerator.status == HardwareStatus.MISSING
assert service.publish_inventory(node, inventory(3, hostname="renamed-host")).accepted
session.refresh(first_accelerator)
assert first_accelerator.id == first_id
assert first_accelerator.status == HardwareStatus.ACTIVE
assert session.query(ComputeNode).count() == 1
assert session.query(Accelerator).count() == 1
assert node.key == "remote-node" and node.hostname == "renamed-host"
def test_older_telemetry_cannot_overwrite_latest(session: Session) -> None:
service = NodeAgentService(session, settings())
_created, response = enroll(service)
_credential, node = service.authenticate(f"Bearer {response.node_credential}")
service.publish_inventory(node, inventory(1))
newer = datetime.now(UTC)
assert service.publish_telemetry(node, telemetry(newer, 1)).accepted
older = telemetry(newer - timedelta(seconds=2), 2)
assert not service.publish_telemetry(node, older).accepted
session.refresh(node)
assert node.telemetry_sequence == 2
assert node.last_telemetry_received_at is not None
assert not service.publish_telemetry(node, older).accepted
assert service.publish_telemetry(node, telemetry(newer + timedelta(seconds=1), 3)).accepted
future = telemetry(newer + timedelta(hours=1), 4)
assert not service.publish_telemetry(node, future).accepted
session.refresh(node)
assert node.telemetry_sequence == 4
assert node.last_connection_error == "telemetry observation exceeds clock-skew limit"
def test_liveness_transitions_and_return_online_are_audited(session: Session) -> None:
service = NodeAgentService(session, settings())
_created, response = enroll(service)
_credential, node = service.authenticate(f"Bearer {response.node_credential}")
baseline = datetime.now(UTC)
node.last_heartbeat_at = baseline
session.commit()
service.evaluate_liveness(baseline + timedelta(seconds=11))
assert node.liveness_state == NodeLiveness.STALE
service.evaluate_liveness(baseline + timedelta(seconds=21))
assert node.liveness_state == NodeLiveness.OFFLINE
service.heartbeat(
node,
HeartbeatRequest(
identity_key="remote-node", metadata=metadata(), observed_at=datetime.now(UTC)
),
)
assert node.liveness_state == NodeLiveness.ONLINE
actions = list(session.scalars(select(AuditEvent.action)))
assert {"NODE_BECAME_STALE", "NODE_BECAME_OFFLINE", "NODE_RETURNED_ONLINE"} <= set(actions)
def test_a_single_use_enrollment_token_can_never_mint_two_node_identities(
session: Session,
) -> None:
"""M15 node-recovery rehearsal regression.
Two agent threads racing to enrol both passed the `used_at` read before either wrote it,
so one token produced two compute nodes for the same hardware. The claim is now atomic.
"""
service = NodeAgentService(session, settings())
created = service.create_enrollment(EnrollmentTokenCreate(display_name="Recovery rehearsal"))
first = service.enroll(
EnrollmentRequest(
enrollment_token=created.enrollment_token,
identity_key="recovery-node",
identity_source="persisted_uuid",
hostname="recovery-host",
display_name="recovery-host",
metadata=metadata(),
)
)
assert first.node_credential
with pytest.raises(AgentAuthenticationError, match="already used"):
service.enroll(
EnrollmentRequest(
enrollment_token=created.enrollment_token,
identity_key="recovery-node-second-identity",
identity_source="persisted_uuid",
hostname="recovery-host",
display_name="recovery-host",
metadata=metadata(),
)
)
nodes = list(session.scalars(select(ComputeNode)))
assert len(nodes) == 1
credentials = list(session.scalars(select(NodeCredential)))
assert len([item for item in credentials if item.revoked_at is None]) == 1
enrollment = session.scalar(select(NodeEnrollment))
assert enrollment is not None
assert enrollment.used_at is not None
assert enrollment.enrolled_node_id == nodes[0].id
def test_re_enrollment_reuses_the_persisted_node_identity_and_revokes_the_old_credential(
session: Session,
) -> None:
"""A recovered node keeps its identity; the credential it lost stops working."""
service = NodeAgentService(session, settings())
_first_token, first = enroll(service, identity="gpu_node-hardware")
node = session.scalar(select(ComputeNode))
assert node is not None
original_node_id = node.id
original_credential = first.node_credential
second_token = service.create_enrollment(EnrollmentTokenCreate(display_name="Re-enrollment"))
second = service.enroll(
EnrollmentRequest(
enrollment_token=second_token.enrollment_token,
identity_key="gpu_node-hardware",
identity_source="persisted_uuid",
hostname="server01",
display_name="server01",
metadata=metadata(),
)
)
assert len(list(session.scalars(select(ComputeNode)))) == 1
assert session.scalar(select(ComputeNode)).id == original_node_id
assert second.node_credential != original_credential
credentials = list(session.scalars(select(NodeCredential)))
assert len(credentials) == 2
assert len([item for item in credentials if item.revoked_at is None]) == 1
with pytest.raises(AgentAuthenticationError):
service.authenticate(f"Bearer {original_credential}")
_credential, authenticated = service.authenticate(f"Bearer {second.node_credential}")
assert authenticated.id == original_node_id