Files
ModelForge/backend/tests/test_acquisition.py
T

498 lines
17 KiB
Python

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