Files
geointel/backend/tests/test_model_asset_catalog.py
Jens faeb58ef6d
GeoIntel release gates / Compile, test, contracts and builds (push) Successful in 1m49s
GeoIntel release gates / Python and npm vulnerability policy (push) Successful in 21s
GeoIntel release gates / Production AI image, SBOM and container scan (push) Successful in 5m39s
GeoIntel release gates / Deploy exact gated revision to Unraid (push) Failing after 58m43s
Initial public release
2026-08-31 21:56:53 +02:00

356 lines
13 KiB
Python

from __future__ import annotations
from hashlib import sha256
import json
from pathlib import Path
from uuid import uuid4
import pytest
from fastapi.testclient import TestClient
from app.core.config import Settings
from app.core.errors import AppError
from app.main import app
from app.models import (
AnalysisRun,
Dataset,
DatasetVersion,
Detection,
Job,
Project,
SourceRegistry,
SourceSnapshot,
)
from app.services.detection_service import DetectionService
from app.services.model_asset_catalog_service import ModelAssetCatalogService
from app.services.runtime_model_provenance_service import RuntimeModelProvenanceService
from app.services.tile_manifest_service import TileManifestService
class FakeSession:
def __init__(self, objects=None) -> None:
self.objects = objects or {}
self.added = []
self.commits = 0
self.refreshes = []
def get(self, model, item_id):
return self.objects.get((model, item_id))
def add(self, item) -> None:
self.added.append(item)
if getattr(item, "id", None) is not None:
self.objects[(item.__class__, item.id)] = item
def commit(self) -> None:
self.commits += 1
def refresh(self, item) -> None:
self.refreshes.append(item)
class MockYoloAdapter:
def __init__(self, settings: Settings) -> None:
self.settings = settings
@staticmethod
def dependencies_available() -> bool:
return True
def load_model(self, model_path: Path):
return {"model_path": str(model_path)}
def predict_tiles(self, model, tile_paths, confidence_threshold: float) -> list[list[dict]]:
# The service batches tiles; this double still answers per tile.
return [self.predict_tile(model, tile_path, confidence_threshold) for tile_path in tile_paths]
def predict_tile(self, model, tile_path: Path, confidence_threshold: float) -> list[dict]:
assert model["model_path"].endswith("building-detector.pt")
return [
{
"class_name": "building",
"confidence": 0.9,
"bbox": [10.0, 20.0, 30.0, 40.0],
"properties": {"adapter": "mock"},
}
]
def _project_and_raster_dataset():
project_id = uuid4()
dataset_id = uuid4()
source_registry_id = uuid4()
source_snapshot_id = uuid4()
checksum = "a" * 64
project = Project(id=project_id, name="Geel")
source_registry = SourceRegistry(
id=source_registry_id,
source_key="test-derived-raster",
display_name="Governed test-derived raster",
classification="derived",
authority_name="GeoIntel test fixture",
usage_policy_json={"ground_truth_allowed": False},
)
source_snapshot = SourceSnapshot(
id=source_snapshot_id,
source_registry_id=source_registry_id,
snapshot_key="test-derived-raster-v1",
checksum_sha256=checksum,
freshness_status="current",
ingest_status="ingested",
)
dataset = Dataset(
id=dataset_id,
project_id=project_id,
name="source.tif",
dataset_type="raster",
source="test-derived-raster",
source_name="test-derived-raster",
storage_path="storage/uploads/source.tif",
checksum_sha256=checksum,
crs="EPSG:4326",
bounds_json={"min_x": 4.0, "min_y": 51.0, "max_x": 5.0, "max_y": 52.0},
source_registry_id=source_registry_id,
source_snapshot_id=source_snapshot_id,
data_contract_key="geointel.raster.geotiff",
data_contract_version="1.0.0",
validation_status="passed",
provenance_status="complete",
lineage_status="not_applicable",
quarantine_status="not_quarantined",
status="ready",
)
dataset.source_registry = source_registry
dataset.source_snapshot = source_snapshot
dataset.versions.append(
DatasetVersion(id=uuid4(), dataset_id=dataset_id, version=1, checksum_sha256=checksum)
)
db = FakeSession(objects={(Project, project_id): project, (Dataset, dataset_id): dataset})
return db, project_id, dataset_id
def _manifest(tmp_path: Path, db: FakeSession, dataset: Dataset) -> Path:
tile_path = tmp_path / "tile_0000.tif"
tile_path.write_bytes(b"tile")
binding = TileManifestService.dataset_binding(db, dataset)
manifest_path = tmp_path / "manifest.json"
manifest_path.write_text(
json.dumps(
{
**binding,
"tile_set_id": "tiles-fixture",
"count": 1,
"crs": "EPSG:4326",
"bounds": [4.0, 51.0, 5.0, 52.0],
"tiles": [
{
"path": str(tile_path),
"pixel_window": [0, 0, 100, 100],
"bounds": [4.0, 51.0, 5.0, 52.0],
"transform": [4.0, 0.01, 0.0, 52.0, 0.0, -0.01],
"crs": "EPSG:4326",
"index": 0,
**TileManifestService.tile_integrity(tile_path),
}
],
}
),
encoding="utf-8",
)
return manifest_path
def _write_model_sidecar(
model_path: Path,
settings: Settings,
*,
db: FakeSession | None = None,
) -> None:
model_sha256 = sha256(model_path.read_bytes()).hexdigest()
source_registry_id = uuid4()
source_snapshot_id = uuid4()
source_version = settings.yolo_model_version or "test-v1"
if db is not None:
source_registry = SourceRegistry(
id=source_registry_id,
source_key="model",
display_name="Governed test model artifact",
classification="experimental",
authority_name="GeoIntel test fixture",
freshness_status="current",
ingest_status="configured",
)
source_snapshot = SourceSnapshot(
id=source_snapshot_id,
source_registry_id=source_registry_id,
snapshot_key=f"model-{source_version}",
source_version=source_version,
checksum_sha256=model_sha256,
freshness_status="current",
ingest_status="ingested",
)
db.objects[(SourceRegistry, source_registry_id)] = source_registry
db.objects[(SourceSnapshot, source_snapshot_id)] = source_snapshot
payload = {
"schema_version": RuntimeModelProvenanceService.MANIFEST_SCHEMA_VERSION,
"data_contract": {"key": "geointel.model.pytorch", "version": "1.0.0"},
"model": {
"model_id": settings.yolo_model_id,
"task_type": "object_detection",
"sha256": model_sha256,
"model_format": "pytorch",
"framework": "ultralytics/pytorch",
"class_mapping": {"0": "building"},
"source_version": source_version,
},
"source": {
"source_registry_id": str(source_registry_id),
"source_snapshot_id": str(source_snapshot_id),
"source_registry_key": "model",
"source_snapshot_checksum_sha256": model_sha256,
},
"lineage": {
"upstream_asset_ids": ["test-training-corpus"],
"upstream_checksums_sha256": ["a" * 64],
"transformations": [
{"name": "test-training", "version": "1.0.0", "checksum_sha256": "b" * 64}
],
},
"metadata": {"training_manifest_sha256": "c" * 64},
"imported_at": "2026-08-01T10:00:00+00:00",
}
payload["metadata"]["runtime_manifest_sha256"] = RuntimeModelProvenanceService.manifest_self_checksum(payload)
RuntimeModelProvenanceService.manifest_path_for_model(model_path).write_text(
json.dumps(payload, sort_keys=True),
encoding="utf-8",
)
def test_model_asset_catalog_lists_supported_local_model_files(tmp_path: Path) -> None:
model_file = tmp_path / "building-detector.pt"
model_file.write_bytes(b"local model")
ignored_file = tmp_path / "notes.txt"
ignored_file.write_text("ignore me", encoding="utf-8")
settings = Settings(yolo_models_dir=str(tmp_path), yolo_model_path=str(model_file), yolo_enabled=True)
response = ModelAssetCatalogService.list_assets(settings=settings)
assert response.total == 1
asset = response.items[0]
assert asset.model_asset_id == "building-detector-pt"
assert asset.filename == "building-detector.pt"
assert asset.display_name == "building-detector"
assert asset.model_path == str(model_file)
assert asset.size_bytes == len(b"local model")
assert len(asset.sha256) == 64
assert asset.active is True
assert asset.runtime_available is True
assert asset.runtime_status == "active"
assert asset.governed_validation_status == "not_verified_by_catalog"
assert asset.promotion_status == "not_verified_by_catalog"
assert asset.status == "runtime_active"
assert asset.will_download_models is False
def test_model_asset_catalog_resolves_known_asset(tmp_path: Path) -> None:
model_file = tmp_path / "building-detector.pt"
model_file.write_bytes(b"local model")
settings = Settings(yolo_models_dir=str(tmp_path), yolo_enabled=True)
asset = ModelAssetCatalogService.resolve_asset("building-detector-pt", settings=settings)
assert asset.filename == "building-detector.pt"
assert asset.model_path == str(model_file)
def test_model_asset_catalog_only_exposes_explicit_active_asset_in_runtime(tmp_path: Path) -> None:
active_file = tmp_path / "approved-building-detector.pt"
active_file.write_bytes(b"approved")
(tmp_path / "training-smoke.pt").write_bytes(b"experiment")
(tmp_path / "partial-checkpoint.pt").write_bytes(b"partial")
settings = Settings(
yolo_models_dir=str(tmp_path),
yolo_model_path=str(active_file),
yolo_enabled=True,
)
response = ModelAssetCatalogService.list_assets(settings=settings)
assert response.total == 1
assert response.items[0].filename == active_file.name
assert response.items[0].active is True
assert response.items[0].runtime_status == "active"
assert response.items[0].governed_validation_status == "not_verified_by_catalog"
assert response.items[0].promotion_status == "not_verified_by_catalog"
assert response.items[0].status == "runtime_active"
def test_model_asset_catalog_rejects_unknown_asset(tmp_path: Path) -> None:
settings = Settings(yolo_models_dir=str(tmp_path), yolo_enabled=True)
with pytest.raises(AppError) as exc_info:
ModelAssetCatalogService.resolve_asset("missing-model", settings=settings)
assert exc_info.value.code == "DETECTION_MODEL_ASSET_NOT_FOUND"
assert exc_info.value.status_code == 404
def test_model_assets_api_returns_canonical_envelope(monkeypatch, tmp_path: Path) -> None:
model_file = tmp_path / "building-detector.pt"
model_file.write_bytes(b"local model")
monkeypatch.setenv("YOLO_MODELS_DIR", str(tmp_path))
monkeypatch.setenv("YOLO_MODEL_PATH", str(model_file))
response = TestClient(app).get("/api/v1/detection/model-assets")
assert response.status_code == 200
payload = response.json()
assert set(payload) == {"data"}
assert payload["data"]["total"] == 1
assert payload["data"]["items"][0]["model_asset_id"] == "building-detector-pt"
assert payload["data"]["items"][0]["active"] is True
assert payload["data"]["items"][0]["runtime_status"] == "active"
assert payload["data"]["items"][0]["governed_validation_status"] == "not_verified_by_catalog"
assert payload["data"]["items"][0]["promotion_status"] == "not_verified_by_catalog"
assert payload["data"]["items"][0]["will_download_models"] is False
def test_detection_run_persists_selected_model_asset_parameters(tmp_path, monkeypatch: Path) -> None:
# A manifest written into tmp_path is only a governed artifact if
# tmp_path is the storage root.
monkeypatch.setenv("STORAGE_ROOT", str(tmp_path))
model_file = tmp_path / "building-detector.pt"
model_file.write_bytes(b"local model")
db, project_id, dataset_id = _project_and_raster_dataset()
settings = Settings(
yolo_enabled=True,
yolo_model_path=str(tmp_path / "default.pt"),
yolo_models_dir=str(tmp_path),
yolo_max_tiles=4,
)
_write_model_sidecar(model_file, settings, db=db)
result = DetectionService.run_detection(
db=db,
project_id=project_id,
dataset_id=dataset_id,
model_id="yolo-configured",
model_asset_id="building-detector-pt",
confidence_threshold=0.5,
tile_manifest_path=str(_manifest(tmp_path, db, db.get(Dataset, dataset_id))),
settings=settings,
yolo_adapter_class=MockYoloAdapter,
)
jobs = [item for item in db.added if isinstance(item, Job)]
runs = [item for item in db.added if isinstance(item, AnalysisRun)]
detections = [item for item in db.added if isinstance(item, Detection)]
assert result.status == "success"
assert result.detection_count == 1
assert jobs[0].parameters_json["model_asset_id"] == "building-detector-pt"
assert jobs[0].parameters_json["model_asset_path"] == str(model_file)
assert len(jobs[0].parameters_json["model_asset_sha256"]) == 64
assert runs[0].parameters_json["model_asset_id"] == "building-detector-pt"
assert detections[0].model_name == "yolo-configured"