Initial public ModelForge release
This commit is contained in:
@@ -0,0 +1,345 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user