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()