from __future__ import annotations import hashlib import uuid from datetime import UTC, datetime import pytest from sqlalchemy import create_engine from sqlalchemy.orm import Session from modelforge_api.domain.acquisition import ( AgentJobComplete, AgentJobFailure, AgentJobProgress, CompletedFile, DiscoverySearchRequest, DownloadPlanCreate, ) from modelforge_api.persistence.models import ( ArtifactInspection, ArtifactSetMember, AuditEvent, Base, ComputeNode, Model, StorageRoot, ) from modelforge_api.providers.huggingface import ( HuggingFaceGated, ProviderFile, ProviderSnapshot, ) from modelforge_api.services.acquisition import ( AcquisitionError, AcquisitionService, ProviderBoundaryError, ) from modelforge_api.settings import Settings class FakeProvider: def __init__(self, *, fail: Exception | None = None) -> None: self.fail = fail self.sha = "a" * 40 self.files = ( ProviderFile("config.json", 8, "git-config", None, "json", "configuration"), ProviderFile( "model.safetensors", 12, "lfs-weight", hashlib.sha256(b"safe-weights").hexdigest(), "safetensors", "weights", ), ProviderFile( "pytorch_model.bin", 12, "lfs-pickle", None, "bin", "weights", ("pickle_or_executable_serialization",), ), ProviderFile( "modeling_custom.py", 20, "git-code", None, "python", "repository_code", ("remote_code",), ), ) def search(self, query: str, *, limit: int, sort: str, pipeline_tag: str | None): return [ { "repository_id": "org/safe-model", "resolved_commit_sha": self.sha, "access_state": "public", "pipeline_tag": pipeline_tag or "feature-extraction", "library_name": "transformers", "tags": ["safetensors"], "downloads": 10, "likes": 2, "last_modified": datetime.now(UTC), } ][:limit] def snapshot(self, repository_id: str, revision: str) -> ProviderSnapshot: if self.fail: raise self.fail return ProviderSnapshot( repository_id=repository_id, requested_revision=revision, resolved_commit_sha=self.sha, access_state="public", metadata={"pipeline_tag": "feature-extraction", "tags": ["safetensors"]}, card_metadata={"license": "apache-2.0"}, security_metadata={"upstream_scanner": {"status": "safe"}, "evidence_only": True}, source_updated_at=datetime.now(UTC), files=self.files, ) @pytest.fixture def session() -> Session: engine = create_engine("sqlite+pysqlite:///:memory:") Base.metadata.create_all(engine) with Session(engine) as value: yield value def setup_service(session: Session, *, fail: Exception | None = None): model = Model( key="safe-model", display_name="Safe model", source_type="huggingface", upstream_provider="org", upstream_source="org/safe-model", upstream_metadata={}, local_metadata={"owner": "operator"}, interpretation_metadata={"seed": True}, modalities=[], parameter_metadata={}, license_metadata={"status": "unknown"}, lifecycle="candidate", ) node = ComputeNode( key=f"node-{uuid.uuid4()}", hostname="gpu_node", enabled=True, liveness_state="online", agent_capabilities=["hardware.inventory", "artifact.acquire.v1"], ) session.add_all([model, node]) session.flush() root = StorageRoot( compute_node_id=node.id, name="gpu_node-cache", path="/host/cache/models", agent_path="/data/artifacts/model-registry", status="ready", writable=True, capacity_bytes=10_000, free_bytes=8_000, reserve_bytes=100, reserve_percent=10, ) session.add(root) session.commit() settings = Settings( database_url="sqlite+pysqlite:///:memory:", hf_snapshot_ttl_seconds=3600, ) return AcquisitionService(session, settings, FakeProvider(fail=fail)), model, node, root def test_discovery_marks_upstream_facts_and_local_interpretation(session: Session) -> None: service, model, _node, _root = setup_service(session) result = service.search(DiscoverySearchRequest(query="safe")) assert result[0].matched_model_id == model.id assert result[0].upstream_facts["source"] == "huggingface_hub.list_models" assert result[0].local_interpretation["approval"] == "not_evaluated" def test_refresh_pins_exact_revision_and_prefers_safetensors_without_approval( session: Session, ) -> None: service, model, _node, _root = setup_service(session) snapshot = service.refresh_model(model.id, "main") session.refresh(model) revisions = model.revisions assert snapshot.resolved_commit_sha == "a" * 40 assert len(revisions) == 1 and revisions[0].immutable_at is not None artifact_sets = service.artifact_sets(revisions[0].id) assert artifact_sets[0].selected_paths == ["config.json", "model.safetensors"] assert "pytorch_model.bin" not in artifact_sets[0].selected_paths assert model.lifecycle == "candidate" assert model.local_metadata == {"owner": "operator"} assert model.license_metadata["status"] == "captured_unreviewed" def test_refresh_is_repeatable_without_duplicate_revision(session: Session) -> None: service, model, _node, _root = setup_service(session) first = service.refresh_model(model.id, "main") second = service.refresh_model(model.id, "main") assert first.id != second.id assert len(model.revisions) == 1 assert len(service.artifact_sets(model.revisions[0].id)) == 1 def test_moving_mutable_ref_creates_new_revision_without_changing_original( session: Session, ) -> None: service, model, _node, _root = setup_service(session) first = service.refresh_model(model.id, "main") provider = service.provider assert isinstance(provider, FakeProvider) provider.sha = "b" * 40 second = service.refresh_model(model.id, "main") assert first.resolved_commit_sha == "a" * 40 assert second.resolved_commit_sha == "b" * 40 assert {revision.resolved_commit_sha for revision in model.revisions} == { "a" * 40, "b" * 40, } def test_gated_error_is_explicit_and_does_not_mutate_candidate(session: Session) -> None: service, model, _node, _root = setup_service(session, fail=HuggingFaceGated("gated")) with pytest.raises(ProviderBoundaryError) as error: service.refresh_model(model.id, "main") assert error.value.details == {"provider_code": "repository_gated", "access_state": "gated"} assert model.revisions == [] def test_immutable_plan_has_capacity_and_node_preflights(session: Session) -> None: service, model, node, root = setup_service(session) service.refresh_model(model.id, "main") artifact_set = service.artifact_sets(model.revisions[0].id)[0] plan = service.create_plan( DownloadPlanCreate( artifact_set_id=artifact_set.id, compute_node_id=node.id, storage_root_id=root.id, ) ) assert plan.resolved_commit_sha == "a" * 40 assert plan.preflight["capacity"]["allowed"] is True assert plan.preflight["trust_remote_code"] is False assert plan.status == "planned" assert len(plan.files) == 2 assert ( service.create_plan( DownloadPlanCreate( artifact_set_id=artifact_set.id, compute_node_id=node.id, storage_root_id=root.id, ) ).id == plan.id ) @pytest.mark.parametrize("mutation", ["offline", "no-capability", "no-agent-path", "no-capacity"]) def test_planning_fails_closed_for_invalid_target(session: Session, mutation: str) -> None: service, model, node, root = setup_service(session) service.refresh_model(model.id, "main") artifact_set = service.artifact_sets(model.revisions[0].id)[0] if mutation == "offline": node.liveness_state = "offline" elif mutation == "no-capability": node.agent_capabilities = [] elif mutation == "no-agent-path": root.agent_path = None else: root.free_bytes = None session.commit() with pytest.raises(AcquisitionError): service.create_plan( DownloadPlanCreate( artifact_set_id=artifact_set.id, compute_node_id=node.id, storage_root_id=root.id, ) ) def test_job_lease_progress_completion_registers_verified_content_not_approval( session: Session, ) -> None: service, model, node, root = setup_service(session) service.refresh_model(model.id, "main") artifact_set = service.artifact_sets(model.revisions[0].id)[0] plan = service.create_plan( DownloadPlanCreate( artifact_set_id=artifact_set.id, compute_node_id=node.id, storage_root_id=root.id ) ) with pytest.raises(AcquisitionError, match="explicit approval"): service.execute_plan(plan.id) plan = service.approve_plan(plan.id) assert plan.status == "ready" job = service.execute_plan(plan.id) assert len(service.repo.job(job.id).idempotency_key) == 64 # type: ignore[union-attr] lease = service.claim_next(node) assert lease is not None and lease.job_id == job.id control = service.progress( job.id, node, AgentJobProgress( lease_token=lease.lease_token, status="verifying", progress_bytes=0, current_file="config.json", quarantine_relative_path=f".quarantine/{job.id}", ), ) assert control.accepted and not control.cancel_requested payloads = {"config.json": b"12345678", "model.safetensors": b"safe-weights"} completed = service.complete( job.id, node, AgentJobComplete( lease_token=lease.lease_token, promoted_relative_path="repositories/org--safe-model/" + "a" * 40, capacity_observation={"free_bytes": 7000}, files=[ CompletedFile( path=file.path, relative_path="repositories/org--safe-model/" + "a" * 40 + "/" + file.path, size_bytes=len(payloads[file.path]), sha256=hashlib.sha256(payloads[file.path]).hexdigest(), inspections=[ {"type": "static", "status": "passed", "severity": "info", "evidence": {}} ], ) for file in lease.files ], ), ) assert completed.status == "completed" session.refresh(model) assert model.lifecycle == "candidate" assert all(item.status == "verified" for item in model.revisions[0].artifacts) assert all( item.security_status == "static_checks_passed_unapproved" for item in model.revisions[0].artifacts ) assert session.query(ArtifactSetMember).count() == 2 assert session.query(ArtifactInspection).count() == 2 assert root.free_bytes == 8_000 # incomplete observations never overwrite last-good capacity actions = {item.action for item in session.query(AuditEvent).all()} assert { "DOWNLOAD_PLAN_APPROVED", "ARTIFACT_DOWNLOAD_QUEUED", "ARTIFACT_DOWNLOAD_STARTED", "ARTIFACT_QUARANTINED", "ARTIFACT_VERIFIED", "ARTIFACT_PROMOTED", } <= actions def test_retry_is_bounded_and_cancel_is_terminal(session: Session) -> None: service, model, node, root = setup_service(session) service.refresh_model(model.id, "main") artifact_set = service.artifact_sets(model.revisions[0].id)[0] plan = service.create_plan( DownloadPlanCreate( artifact_set_id=artifact_set.id, compute_node_id=node.id, storage_root_id=root.id ) ) service.approve_plan(plan.id) job = service.execute_plan(plan.id) for expected in ("queued", "queued", "failed"): lease = service.claim_next(node) assert lease is not None failed = service.fail( job.id, node, AgentJobFailure( lease_token=lease.lease_token, error_code="network", error_message="temporary", retryable=True, ), ) assert failed.status == expected assert failed.attempt_count == 3 retried = service.retry(job.id) assert retried.status == "queued" assert retried.attempt_count == 3 assert retried.progress_bytes == 0 assert retried.result["operator_retry"] is True assert service.claim_next(node) is not None assert service.job(job.id).attempt_count == 4 def test_stale_plan_never_executes(session: Session, monkeypatch: pytest.MonkeyPatch) -> None: service, model, node, root = setup_service(session) service.refresh_model(model.id, "main") artifact_set = service.artifact_sets(model.revisions[0].id)[0] plan = service.create_plan( DownloadPlanCreate( artifact_set_id=artifact_set.id, compute_node_id=node.id, storage_root_id=root.id ) ) service.approve_plan(plan.id) monkeypatch.setattr( "modelforge_api.services.acquisition._utcnow", lambda: plan.expires_at.replace(tzinfo=UTC, year=plan.expires_at.year + 1), ) with pytest.raises(AcquisitionError, match="stale"): service.execute_plan(plan.id) def test_gated_snapshot_requires_configured_credentials_for_planning(session: Session) -> None: service, model, node, root = setup_service(session) snapshot = service.refresh_model(model.id, "main") stored = service.repo.snapshot(snapshot.id) assert stored is not None stored.access_state = "gated" session.commit() artifact_set = service.artifact_sets(model.revisions[0].id)[0] with pytest.raises(AcquisitionError, match="authentication is required"): service.create_plan( DownloadPlanCreate( artifact_set_id=artifact_set.id, compute_node_id=node.id, storage_root_id=root.id, ) ) def test_job_is_claimed_only_by_its_target_node_and_execution_is_idempotent( session: Session, ) -> None: service, model, node, root = setup_service(session) other = ComputeNode( key="other-node", hostname="workstation", enabled=True, liveness_state="online", agent_capabilities=["artifact.acquire.v1"], ) session.add(other) session.commit() service.refresh_model(model.id, "main") artifact_set = service.artifact_sets(model.revisions[0].id)[0] plan = service.create_plan( DownloadPlanCreate( artifact_set_id=artifact_set.id, compute_node_id=node.id, storage_root_id=root.id, ) ) service.approve_plan(plan.id) first = service.execute_plan(plan.id) second = service.execute_plan(plan.id) assert second.id == first.id assert service.claim_next(other) is None assert service.claim_next(node) is not None def test_execution_rechecks_capacity_after_plan_approval(session: Session) -> None: service, model, node, root = setup_service(session) service.refresh_model(model.id, "main") artifact_set = service.artifact_sets(model.revisions[0].id)[0] plan = service.create_plan( DownloadPlanCreate( artifact_set_id=artifact_set.id, compute_node_id=node.id, storage_root_id=root.id, ) ) service.approve_plan(plan.id) root.free_bytes = 100 session.commit() with pytest.raises(AcquisitionError, match="execution capacity preflight failed"): service.execute_plan(plan.id) def test_download_plan_payload_is_immutable(session: Session) -> None: service, model, node, root = setup_service(session) service.refresh_model(model.id, "main") artifact_set = service.artifact_sets(model.revisions[0].id)[0] plan = service.create_plan( DownloadPlanCreate( artifact_set_id=artifact_set.id, compute_node_id=node.id, storage_root_id=root.id, ) ) stored = service.repo.plan(plan.id) assert stored is not None stored.resolved_commit_sha = "b" * 40 with pytest.raises(ValueError, match="immutable approved fields"): session.commit() def test_queued_job_can_be_cancelled_without_claim_or_promotion(session: Session) -> None: service, model, node, root = setup_service(session) service.refresh_model(model.id, "main") artifact_set = service.artifact_sets(model.revisions[0].id)[0] plan = service.create_plan( DownloadPlanCreate( artifact_set_id=artifact_set.id, compute_node_id=node.id, storage_root_id=root.id, ) ) service.approve_plan(plan.id) job = service.execute_plan(plan.id) cancelled = service.cancel(job.id) assert cancelled.status == "cancelled" assert service.claim_next(node) is None