Files
ModelForge/backend/tests/test_registry.py
T

346 lines
13 KiB
Python

from __future__ import annotations
import hashlib
import uuid
from pathlib import Path
import pytest
from fastapi.testclient import TestClient
from pydantic import SecretStr
from sqlalchemy import create_engine, select
from sqlalchemy.orm import Session
from sqlalchemy.pool import StaticPool
from modelforge_api.api.routes.registry import get_registry_service
from modelforge_api.domain.enums import ArtifactStatus, StorageRootStatus
from modelforge_api.domain.registry import (
ArtifactCreate,
ArtifactLocationCreate,
DerivedArtifactCreate,
ModelCreate,
ModelUpdate,
RevisionCreate,
StorageRootCreate,
StorageRootObservation,
)
from modelforge_api.main import app
from modelforge_api.persistence.models import AuditEvent, Base, ComputeNode, Model, ModelRevision
from modelforge_api.services.manifest_registry import ManifestRegistry
from modelforge_api.services.registry import (
RegistryConflict,
RegistryService,
seed_candidate_registry,
)
from modelforge_api.settings import Settings, get_settings
def _test_operator_credential() -> str:
return "registry-test-operator"
def _operator_headers() -> dict[str, str]:
return {"X-ModelForge-Admin-Token": _test_operator_credential()}
@pytest.fixture
def session() -> Session:
engine = create_engine("sqlite+pysqlite:///:memory:")
Base.metadata.create_all(engine)
with Session(engine) as value:
yield value
def create_model(service: RegistryService, key: str = "registry-test"):
return service.create_model(
ModelCreate(
key=key,
display_name="Registry Test",
description="locally governed metadata",
upstream_provider="example",
upstream_source=f"example/{key}",
upstream_metadata={"repository_id": f"example/{key}"},
local_metadata={"owner": "modelops"},
interpretation_metadata={"intended_capabilities": ["assistant.general"]},
license_metadata={"status": "unknown", "spdx_id": None},
)
)
def create_revision(service: RegistryService, model_id: uuid.UUID, suffix: str = "a"):
return service.create_revision(
model_id,
RevisionCreate(
upstream_revision="main",
resolved_commit_sha=suffix * 40,
metadata_snapshot={"source": "operator"},
),
)
def create_node_and_root(service: RegistryService, session: Session, path: Path):
node = ComputeNode(key=f"node-{uuid.uuid4()}", hostname="test-node")
session.add(node)
session.commit()
return service.create_storage_root(
StorageRootCreate(
compute_node_id=node.id,
name="model-cache",
path=str(path),
reserve_bytes=10,
reserve_percent=10,
)
)
def test_seed_is_idempotent_and_preserves_local_admin_metadata(session: Session) -> None:
manifests = ManifestRegistry()
assert seed_candidate_registry(session, manifests) == 15
model = session.scalar(select(Model).where(Model.key == "qwen-general"))
assert model is not None
model.local_metadata = {"owner": "Jens", "review_status": "reviewed"}
session.commit()
assert seed_candidate_registry(session, manifests) == 0
session.refresh(model)
assert model.local_metadata == {"owner": "Jens", "review_status": "reviewed"}
assert session.query(Model).count() == 15
assert model.license_metadata == {"status": "unknown", "spdx_id": None, "source": None}
def test_model_crud_deprecation_and_audit(session: Session) -> None:
service = RegistryService(session)
model = create_model(service)
updated = service.update_model(
model.id,
ModelUpdate(display_name="Renamed", local_metadata={"owner": "platform"}),
)
assert updated.display_name == "Renamed"
assert updated.local_metadata == {"owner": "platform"}
deprecated = service.deprecate_model(model.id)
assert deprecated.lifecycle == "deprecated" and deprecated.deprecated_at is not None
archived = service.archive_model(model.id)
assert archived.lifecycle == "archived"
service.delete("model", model.id)
assert session.get(Model, model.id) is None
actions = set(session.scalars(select(AuditEvent.action)))
assert {
"MODEL_CREATED",
"MODEL_UPDATED",
"MODEL_DEPRECATED",
"MODEL_ARCHIVED",
"MODEL_DELETED",
} <= actions
def test_exact_revision_is_immutable_and_duplicate_is_rejected(session: Session) -> None:
service = RegistryService(session)
model = create_model(service)
revision = create_revision(service, model.id)
assert revision.immutable_at is not None
with pytest.raises(RegistryConflict, match="already registered"):
create_revision(service, model.id)
entity = session.get(ModelRevision, revision.id)
assert entity is not None
entity.resolved_commit_sha = "b" * 40
with pytest.raises(ValueError, match="immutable approved fields"):
session.commit()
session.rollback()
def test_streaming_hash_verification_good_corrupt_and_missing(
session: Session, tmp_path: Path
) -> None:
service = RegistryService(session)
model = create_model(service)
revision = create_revision(service, model.id)
root = create_node_and_root(service, session, tmp_path)
payload = b"safe static artifact bytes"
target = tmp_path / "weights.bin"
target.write_bytes(payload)
digest = hashlib.sha256(payload).hexdigest()
artifact = service.create_artifact(
revision.id,
ArtifactCreate(
filename="weights.bin",
artifact_type="weights",
serialization_format="safetensors",
sha256=digest,
size_bytes=len(payload),
status=ArtifactStatus.LOCAL,
locations=[
ArtifactLocationCreate(
storage_root_id=root.id,
relative_path="weights.bin",
status=ArtifactStatus.LOCAL,
)
],
),
)
location_id = artifact.locations[0].id
verified = service.verify_artifact(artifact.id, location_id)
assert verified.status is ArtifactStatus.VERIFIED
assert (
service.repo.artifact(artifact.id).verification_details["inspection"]
== "streaming_sha256_only"
) # type: ignore[union-attr]
target.write_bytes(b"tampered")
assert service.verify_artifact(artifact.id, location_id).status is ArtifactStatus.CORRUPT
target.unlink()
assert service.verify_artifact(artifact.id, location_id).status is ArtifactStatus.MISSING
def test_multiple_locations_are_not_artifact_identity(session: Session, tmp_path: Path) -> None:
service = RegistryService(session)
model = create_model(service)
revision = create_revision(service, model.id)
first_root = create_node_and_root(service, session, tmp_path / "one")
second_root = create_node_and_root(service, session, tmp_path / "two")
artifact = service.create_artifact(
revision.id,
ArtifactCreate(
filename="model.safetensors",
artifact_type="weights",
serialization_format="safetensors",
sha256="c" * 64,
size_bytes=42,
locations=[
ArtifactLocationCreate(storage_root_id=first_root.id, relative_path="a/model.bin"),
ArtifactLocationCreate(storage_root_id=second_root.id, relative_path="b/model.bin"),
],
),
)
assert len(artifact.locations) == 2
assert {item.artifact_id for item in artifact.locations} == {artifact.id}
def test_multi_source_derived_lineage_and_dependency_safe_delete(
session: Session,
) -> None:
service = RegistryService(session)
model = create_model(service)
revision = create_revision(service, model.id)
sources = []
for digest, name in (("d" * 64, "weights.bin"), ("e" * 64, "tokenizer.json")):
sources.append(
service.create_artifact(
revision.id,
ArtifactCreate(
filename=name,
artifact_type="source",
serialization_format="raw",
sha256=digest,
size_bytes=10,
),
)
)
derived = service.create_derived(
DerivedArtifactCreate(
revision_id=revision.id,
source_artifact_ids=[item.id for item in sources],
filename="bundle.gguf",
artifact_type="weights",
sha256="f" * 64,
size_bytes=20,
transformation_type="format_conversion",
tool="converter",
tool_version="1.2.3",
configuration={"precision": "fp16"},
environment_snapshot={"container_digest": "sha256:" + "1" * 64},
)
)
assert [item.sha256 for item in derived.sources] == ["d" * 64, "e" * 64]
with pytest.raises(RegistryConflict) as blocked:
service.delete("model_artifact", sources[0].id)
assert blocked.value.details["dependencies"][0]["resource_type"] == "derived_artifact"
with pytest.raises(RegistryConflict):
service.delete("model_revision", revision.id)
actions = list(session.scalars(select(AuditEvent.action)))
assert "DERIVED_ARTIFACT_REGISTERED" in actions
assert "REGISTRY_DELETE_BLOCKED" in actions
def test_storage_capacity_guards_unknown_read_only_and_reserve(
session: Session, tmp_path: Path
) -> None:
service = RegistryService(session)
root = create_node_and_root(service, session, tmp_path)
unknown = service.check_capacity(root.id, 1)
assert not unknown.allowed and unknown.status is StorageRootStatus.UNKNOWN
observed = service.observe_storage_root(
root.id,
StorageRootObservation(
writable=False,
capacity_bytes=1000,
free_bytes=500,
details={"probe": "remote-node-agent"},
),
)
assert observed.status is StorageRootStatus.READ_ONLY
service.observe_storage_root(
root.id,
StorageRootObservation(writable=True, capacity_bytes=1000, free_bytes=500),
)
assert service.check_capacity(root.id, 400).allowed
blocked = service.check_capacity(root.id, 401)
assert not blocked.allowed and blocked.usable_bytes == 400
def test_typed_paginated_api_and_structured_delete_conflict() -> None:
engine = create_engine(
"sqlite+pysqlite:///:memory:",
connect_args={"check_same_thread": False},
poolclass=StaticPool,
)
Base.metadata.create_all(engine)
with Session(engine) as session:
service = RegistryService(session)
settings = Settings(
_env_file=None,
operator_api_key=SecretStr(_test_operator_credential()),
)
app.dependency_overrides[get_registry_service] = lambda: service
app.dependency_overrides[get_settings] = lambda: settings
try:
with TestClient(app) as client:
created = client.post(
"/api/v1/models",
headers=_operator_headers(),
json={
"key": "api-model",
"display_name": "API Model",
"source_type": "custom",
"upstream_provider": "internal",
"upstream_source": "internal/api-model",
"upstream_metadata": {"verification": "unverified"},
"local_metadata": {"owner": "test"},
"interpretation_metadata": {},
"modalities": [],
"parameter_metadata": {},
"license_metadata": {"status": "unknown"},
"lifecycle": "candidate",
},
)
assert created.status_code == 201
model_id = created.json()["id"]
page = client.get(
"/api/v1/models?page=1&page_size=1&search=API",
headers=_operator_headers(),
).json()
assert page["total"] == 1 and page["items"][0]["id"] == model_id
revision = client.post(
f"/api/v1/models/{model_id}/revisions",
headers=_operator_headers(),
json={
"upstream_revision": "release",
"resolved_commit_sha": "a" * 40,
"metadata_snapshot": {},
},
)
assert revision.status_code == 201
conflict = client.delete(f"/api/v1/models/{model_id}", headers=_operator_headers())
assert conflict.status_code == 409
details = conflict.json()["error"]["details"]["dependencies"]
assert details[0]["resource_type"] == "model_revision"
assert details[0]["relation"] == "revision"
finally:
app.dependency_overrides.clear()