346 lines
13 KiB
Python
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()
|