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