498 lines
17 KiB
Python
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
|