Add local model asset catalog
This commit is contained in:
@@ -125,6 +125,7 @@ bash scripts/live_migration_smoke.sh
|
||||
- detection service boundary for creating jobs, analysis runs and dependency-aware unavailable responses.
|
||||
- Added detection endpoints:
|
||||
- `GET /api/v1/detection/models`
|
||||
- `GET /api/v1/detection/model-assets`
|
||||
- `POST /api/v1/detection/run`
|
||||
- `GET /api/v1/detection/runs/{analysis_run_id}`
|
||||
- `GET /api/v1/detection/runs/{analysis_run_id}/detections`
|
||||
@@ -239,6 +240,7 @@ Configured YOLO requires:
|
||||
|
||||
```bash
|
||||
YOLO_ENABLED=true
|
||||
YOLO_MODELS_DIR=/absolute/path/to/models
|
||||
YOLO_MODEL_PATH=/absolute/path/to/local-model.pt
|
||||
```
|
||||
|
||||
@@ -260,6 +262,7 @@ directory, mounted as `/app/models` by default:
|
||||
```bash
|
||||
GEOINTEL_MODELS_PATH=/mnt/user/appdata/geointel/models
|
||||
YOLO_ENABLED=true
|
||||
YOLO_MODELS_DIR=/app/models
|
||||
YOLO_MODEL_PATH=/app/models/local-model.pt
|
||||
```
|
||||
|
||||
@@ -275,6 +278,19 @@ python scripts/configure_yolo_model.py \
|
||||
The smoke loads only the supplied local model file, does not run inference and
|
||||
does not download weights.
|
||||
|
||||
The backend also exposes a read-only model asset catalog for the mounted model
|
||||
directory:
|
||||
|
||||
```bash
|
||||
curl http://localhost:1202/api/v1/detection/model-assets
|
||||
```
|
||||
|
||||
The catalog lists local `.pt`, `.onnx` and `.engine` files with size, SHA-256
|
||||
and active-model status. Detection runs may submit `model_asset_id` with
|
||||
`model_id="yolo-configured"` to use a cataloged local model for that run. The
|
||||
backend resolves the ID to a file inside `YOLO_MODELS_DIR`; browser clients do
|
||||
not send arbitrary model paths.
|
||||
|
||||
Configured YOLO inference uses raster tile artifacts from the existing tile
|
||||
manifest flow. Single-band or otherwise non-RGB tile images are converted to a
|
||||
temporary RGB prediction image before inference; georeferencing still comes
|
||||
@@ -286,6 +302,7 @@ Optional tuning:
|
||||
YOLO_MODEL_ID=yolo-configured
|
||||
YOLO_MODEL_DISPLAY_NAME="Configured YOLO detector"
|
||||
YOLO_MODEL_VERSION=local-v1
|
||||
YOLO_MODELS_DIR=/app/models
|
||||
YOLO_CONFIG_DIR=/app/storage/ultralytics
|
||||
YOLO_DEVICE=cpu
|
||||
YOLO_IMAGE_SIZE=640
|
||||
|
||||
@@ -8,6 +8,7 @@ from sqlalchemy.orm import Session
|
||||
from app.db.session import get_db
|
||||
from app.schemas import DetectionQaRequest, DetectionRunRequest
|
||||
from app.services.detection_service import DetectionService
|
||||
from app.services.model_asset_catalog_service import ModelAssetCatalogService
|
||||
from app.services.model_registry_service import ModelRegistryService
|
||||
from app.services.yolo_preflight_service import YoloPreflightService
|
||||
from app.utils.response import envelope
|
||||
@@ -20,12 +21,22 @@ def list_detection_models() -> dict:
|
||||
return envelope({"models": [model.model_dump() for model in ModelRegistryService.list_model_capabilities()]})
|
||||
|
||||
|
||||
@router.get("/model-assets", response_model=dict)
|
||||
def list_detection_model_assets() -> dict:
|
||||
return envelope(ModelAssetCatalogService.list_assets().model_dump())
|
||||
|
||||
|
||||
@router.get("/yolo/preflight", response_model=dict)
|
||||
def get_yolo_preflight(tile_manifest_path: str | None = None, check_model_load: bool = False) -> dict:
|
||||
def get_yolo_preflight(
|
||||
tile_manifest_path: str | None = None,
|
||||
check_model_load: bool = False,
|
||||
model_asset_id: str | None = None,
|
||||
) -> dict:
|
||||
return envelope(
|
||||
YoloPreflightService.run(
|
||||
tile_manifest_path=tile_manifest_path,
|
||||
check_model_load=check_model_load,
|
||||
model_asset_id=model_asset_id,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -37,6 +48,7 @@ def run_detection(payload: DetectionRunRequest, db: Session = Depends(get_db)) -
|
||||
project_id=payload.project_id,
|
||||
dataset_id=payload.dataset_id,
|
||||
model_id=payload.model_id,
|
||||
model_asset_id=payload.model_asset_id,
|
||||
confidence_threshold=payload.confidence_threshold,
|
||||
class_filter=payload.class_filter,
|
||||
tile_manifest_path=payload.tile_manifest_path,
|
||||
|
||||
@@ -23,6 +23,7 @@ class Settings(BaseSettings):
|
||||
log_level: str = Field(default="INFO", validation_alias="GEOINTEL_LOG_LEVEL")
|
||||
database_statement_timeout_ms: int = Field(default=5_000, validation_alias="DATABASE_STATEMENT_TIMEOUT_MS")
|
||||
yolo_enabled: bool = Field(default=False, validation_alias="YOLO_ENABLED")
|
||||
yolo_models_dir: str = Field(default="/app/models", validation_alias="YOLO_MODELS_DIR")
|
||||
yolo_model_path: str | None = Field(default=None, validation_alias="YOLO_MODEL_PATH")
|
||||
yolo_model_id: str = Field(default="yolo-configured", validation_alias="YOLO_MODEL_ID")
|
||||
yolo_model_display_name: str = Field(default="Configured YOLO detector", validation_alias="YOLO_MODEL_DISPLAY_NAME")
|
||||
|
||||
@@ -15,6 +15,8 @@ from .detection import (
|
||||
DetectionRunRead,
|
||||
DetectionRunRequest,
|
||||
DetectionRunResponse,
|
||||
ModelAssetListResponse,
|
||||
ModelAssetRead,
|
||||
)
|
||||
from .segmentation import (
|
||||
SegmentationListResponse,
|
||||
@@ -105,6 +107,8 @@ __all__ = [
|
||||
"DetectionRunRead",
|
||||
"DetectionRunRequest",
|
||||
"DetectionRunResponse",
|
||||
"ModelAssetListResponse",
|
||||
"ModelAssetRead",
|
||||
"SegmentationListResponse",
|
||||
"SegmentationModelCapability",
|
||||
"SegmentationModelsResponse",
|
||||
|
||||
@@ -22,10 +22,33 @@ class DetectionModelsResponse(BaseModel):
|
||||
models: list[DetectionModelCapability]
|
||||
|
||||
|
||||
class ModelAssetRead(BaseModel):
|
||||
model_asset_id: str
|
||||
filename: str
|
||||
display_name: str
|
||||
model_path: str
|
||||
suffix: str
|
||||
framework: str
|
||||
task_type: str
|
||||
size_bytes: int
|
||||
sha256: str
|
||||
active: bool
|
||||
status: str
|
||||
limitation_message: str
|
||||
will_download_models: bool = False
|
||||
|
||||
|
||||
class ModelAssetListResponse(BaseModel):
|
||||
items: list[ModelAssetRead]
|
||||
total: int
|
||||
model_directory: str
|
||||
|
||||
|
||||
class DetectionRunRequest(BaseModel):
|
||||
project_id: UUID
|
||||
dataset_id: UUID
|
||||
model_id: str
|
||||
model_asset_id: str | None = None
|
||||
confidence_threshold: float = Field(default=0.5, ge=0.0, le=1.0)
|
||||
class_filter: list[str] | None = None
|
||||
tile_manifest_path: str | None = None
|
||||
|
||||
@@ -15,6 +15,7 @@ from app.core.errors import AppError
|
||||
from app.models import AnalysisRun, Dataset, Detection, Job, Project, VectorFeature
|
||||
from app.schemas.detection import DetectionListResponse, DetectionRead, DetectionRunListResponse, DetectionRunRead, DetectionRunResponse
|
||||
from app.services.detection_georeferencing import pixel_bbox_to_epsg4326_polygon
|
||||
from app.services.model_asset_catalog_service import ModelAssetCatalogService
|
||||
from app.services.model_registry_service import ModelRegistryService
|
||||
from app.services.qa_service import QaService
|
||||
from app.services.quality_service import QualityService
|
||||
@@ -33,6 +34,7 @@ class DetectionService:
|
||||
dataset_id: uuid.UUID,
|
||||
model_id: str,
|
||||
confidence_threshold: float,
|
||||
model_asset_id: str | None = None,
|
||||
class_filter: list[str] | None = None,
|
||||
tile_manifest_path: str | None = None,
|
||||
parameters_json: dict[str, Any] | None = None,
|
||||
@@ -55,6 +57,11 @@ class DetectionService:
|
||||
status_code=400,
|
||||
)
|
||||
|
||||
selected_model_asset = None
|
||||
if model_id == resolved_settings.yolo_model_id and model_asset_id:
|
||||
selected_model_asset = ModelAssetCatalogService.resolve_asset(model_asset_id, settings=resolved_settings)
|
||||
resolved_settings = ModelAssetCatalogService.settings_for_asset(resolved_settings, selected_model_asset)
|
||||
|
||||
model = ModelRegistryService.get_model_capability(
|
||||
model_id,
|
||||
settings=resolved_settings,
|
||||
@@ -77,6 +84,9 @@ class DetectionService:
|
||||
|
||||
run_parameters = {
|
||||
"model_id": model.model_id,
|
||||
"model_asset_id": selected_model_asset.model_asset_id if selected_model_asset else None,
|
||||
"model_asset_path": selected_model_asset.model_path if selected_model_asset else None,
|
||||
"model_asset_sha256": selected_model_asset.sha256 if selected_model_asset else None,
|
||||
"confidence_threshold": confidence_threshold,
|
||||
"class_filter": class_filter or [],
|
||||
"tile_manifest_path": tile_manifest_path,
|
||||
|
||||
@@ -0,0 +1,101 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import re
|
||||
from pathlib import Path
|
||||
|
||||
from app.core.config import Settings, get_settings
|
||||
from app.core.errors import AppError
|
||||
from app.schemas.detection import ModelAssetListResponse, ModelAssetRead
|
||||
|
||||
|
||||
class ModelAssetCatalogService:
|
||||
SUPPORTED_SUFFIXES = {
|
||||
".pt": "ultralytics/pytorch",
|
||||
".onnx": "onnx",
|
||||
".engine": "tensorrt",
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def list_assets(settings: Settings | None = None) -> ModelAssetListResponse:
|
||||
resolved_settings = settings or get_settings()
|
||||
model_directory = ModelAssetCatalogService._model_directory(resolved_settings)
|
||||
active_model_path = ModelAssetCatalogService._resolved_file_path(resolved_settings.yolo_model_path)
|
||||
if not model_directory.exists() or not model_directory.is_dir():
|
||||
return ModelAssetListResponse(items=[], total=0, model_directory=str(model_directory))
|
||||
|
||||
items = [
|
||||
ModelAssetCatalogService._asset_from_file(path, active_model_path=active_model_path)
|
||||
for path in sorted(model_directory.iterdir(), key=lambda item: item.name.lower())
|
||||
if path.is_file() and path.suffix.lower() in ModelAssetCatalogService.SUPPORTED_SUFFIXES
|
||||
]
|
||||
return ModelAssetListResponse(items=items, total=len(items), model_directory=str(model_directory))
|
||||
|
||||
@staticmethod
|
||||
def resolve_asset(model_asset_id: str, settings: Settings | None = None) -> ModelAssetRead:
|
||||
normalized = model_asset_id.strip()
|
||||
for asset in ModelAssetCatalogService.list_assets(settings=settings).items:
|
||||
if asset.model_asset_id == normalized:
|
||||
return asset
|
||||
raise AppError(
|
||||
code="DETECTION_MODEL_ASSET_NOT_FOUND",
|
||||
message="Selected local model asset was not found in the configured model directory",
|
||||
details={"model_asset_id": normalized},
|
||||
status_code=404,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def settings_for_asset(settings: Settings, asset: ModelAssetRead) -> Settings:
|
||||
return settings.model_copy(update={"yolo_model_path": asset.model_path})
|
||||
|
||||
@staticmethod
|
||||
def _model_directory(settings: Settings) -> Path:
|
||||
configured_directory = Path(settings.yolo_models_dir).expanduser()
|
||||
if configured_directory.exists() and configured_directory.is_dir():
|
||||
return configured_directory.resolve()
|
||||
active_model_path = ModelAssetCatalogService._resolved_file_path(settings.yolo_model_path)
|
||||
if active_model_path and active_model_path.parent.exists() and active_model_path.parent.is_dir():
|
||||
return active_model_path.parent.resolve()
|
||||
return configured_directory.resolve()
|
||||
|
||||
@staticmethod
|
||||
def _asset_from_file(path: Path, *, active_model_path: Path | None) -> ModelAssetRead:
|
||||
resolved_path = path.resolve()
|
||||
return ModelAssetRead(
|
||||
model_asset_id=ModelAssetCatalogService._asset_id(path),
|
||||
filename=path.name,
|
||||
display_name=path.stem,
|
||||
model_path=str(resolved_path),
|
||||
suffix=path.suffix.lower(),
|
||||
framework=ModelAssetCatalogService.SUPPORTED_SUFFIXES[path.suffix.lower()],
|
||||
task_type="object_detection",
|
||||
size_bytes=path.stat().st_size,
|
||||
sha256=ModelAssetCatalogService._sha256(path),
|
||||
active=active_model_path == resolved_path,
|
||||
status="available",
|
||||
limitation_message="Local runtime model asset. GeoIntel will not download or mutate model weights.",
|
||||
will_download_models=False,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _asset_id(path: Path) -> str:
|
||||
raw = f"{path.stem}-{path.suffix.lower().lstrip('.')}"
|
||||
normalized = re.sub(r"[^a-z0-9]+", "-", raw.lower()).strip("-")
|
||||
return normalized or "model-asset"
|
||||
|
||||
@staticmethod
|
||||
def _sha256(path: Path) -> str:
|
||||
digest = hashlib.sha256()
|
||||
with path.open("rb") as handle:
|
||||
for chunk in iter(lambda: handle.read(1024 * 1024), b""):
|
||||
digest.update(chunk)
|
||||
return digest.hexdigest()
|
||||
|
||||
@staticmethod
|
||||
def _resolved_file_path(raw_path: str | None) -> Path | None:
|
||||
if not raw_path:
|
||||
return None
|
||||
path = Path(raw_path).expanduser()
|
||||
if not path.exists() or not path.is_file():
|
||||
return None
|
||||
return path.resolve()
|
||||
@@ -8,6 +8,7 @@ from typing import Any, Type
|
||||
from app.core.config import Settings, get_settings
|
||||
from app.core.errors import AppError
|
||||
from app.services.detection_service import DetectionService
|
||||
from app.services.model_asset_catalog_service import ModelAssetCatalogService
|
||||
from app.services.yolo_adapter import YoloDetectionAdapter
|
||||
|
||||
|
||||
@@ -20,10 +21,16 @@ class YoloPreflightService:
|
||||
yolo_adapter_class: Type[YoloDetectionAdapter] = YoloDetectionAdapter,
|
||||
assume_dependencies: bool = False,
|
||||
check_model_load: bool = False,
|
||||
model_asset_id: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
resolved_settings = settings or get_settings()
|
||||
selected_asset = None
|
||||
if model_asset_id:
|
||||
selected_asset = ModelAssetCatalogService.resolve_asset(model_asset_id, settings=resolved_settings)
|
||||
resolved_settings = ModelAssetCatalogService.settings_for_asset(resolved_settings, selected_asset)
|
||||
result: dict[str, Any] = {
|
||||
"model_id": resolved_settings.yolo_model_id,
|
||||
"model_asset_id": selected_asset.model_asset_id if selected_asset else None,
|
||||
"model_path": resolved_settings.yolo_model_path,
|
||||
"tile_manifest_path": tile_manifest_path,
|
||||
"status": "not_configured",
|
||||
|
||||
@@ -80,6 +80,7 @@ def test_env_example_uses_runtime_env_names_read_by_backend_and_frontend() -> No
|
||||
|
||||
assert "GEOINTEL_INSTALL_AI=false" in env_example
|
||||
assert "YOLO_ENABLED=false" in env_example
|
||||
assert "YOLO_MODELS_DIR=/app/models" in env_example
|
||||
assert "YOLO_MODEL_PATH=" in env_example
|
||||
assert "YOLO_CONFIG_DIR=./storage/ultralytics" in env_example
|
||||
assert "YOLO_MAX_TILES=100" in env_example
|
||||
@@ -231,6 +232,8 @@ def test_unraid_deploy_passes_ai_build_arg_and_yolo_runtime_env() -> None:
|
||||
|
||||
assert 'YOLO_ENABLED="${YOLO_ENABLED:-false}"' in run_script
|
||||
assert '-e YOLO_ENABLED="$YOLO_ENABLED"' in run_script
|
||||
assert 'YOLO_MODELS_DIR="${YOLO_MODELS_DIR:-/app/models}"' in run_script
|
||||
assert '-e YOLO_MODELS_DIR="$YOLO_MODELS_DIR"' in run_script
|
||||
assert '-e YOLO_MODEL_PATH="$YOLO_MODEL_PATH"' in run_script
|
||||
assert '-e YOLO_MAX_TILES="$YOLO_MAX_TILES"' in run_script
|
||||
assert "-v \"${GEOINTEL_MODELS_PATH}:/app/models\"" in run_script
|
||||
|
||||
@@ -0,0 +1,197 @@
|
||||
from __future__ import annotations
|
||||
|
||||
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, Detection, Job, Project
|
||||
from app.services.detection_service import DetectionService
|
||||
from app.services.model_asset_catalog_service import ModelAssetCatalogService
|
||||
|
||||
|
||||
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_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()
|
||||
project = Project(id=project_id, name="Geel")
|
||||
dataset = Dataset(
|
||||
id=dataset_id,
|
||||
project_id=project_id,
|
||||
name="source.tif",
|
||||
dataset_type="raster",
|
||||
source="user_upload",
|
||||
storage_path="storage/uploads/source.tif",
|
||||
)
|
||||
db = FakeSession(objects={(Project, project_id): project, (Dataset, dataset_id): dataset})
|
||||
return db, project_id, dataset_id
|
||||
|
||||
|
||||
def _manifest(tmp_path: Path) -> Path:
|
||||
tile_path = tmp_path / "tile_0000.tif"
|
||||
tile_path.write_bytes(b"tile")
|
||||
manifest_path = tmp_path / "manifest.json"
|
||||
manifest_path.write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"tile_set_id": "tiles-fixture",
|
||||
"count": 1,
|
||||
"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],
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
}
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
return manifest_path
|
||||
|
||||
|
||||
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.status == "available"
|
||||
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_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]["will_download_models"] is False
|
||||
|
||||
|
||||
def test_detection_run_persists_selected_model_asset_parameters(tmp_path: Path) -> None:
|
||||
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,
|
||||
)
|
||||
|
||||
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)),
|
||||
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"
|
||||
@@ -23,3 +23,29 @@ def test_detection_lab_surfaces_yolo_runtime_preflight() -> None:
|
||||
assert "/api/v1/detection/yolo/preflight" in api
|
||||
assert "interface YoloPreflightResponse" in types
|
||||
assert "yoloPreflight={yoloPreflight}" in app
|
||||
|
||||
|
||||
def test_detection_lab_surfaces_local_model_asset_selection() -> None:
|
||||
lab = (ROOT / "frontend" / "src" / "components" / "detection" / "DetectionLab.tsx").read_text(
|
||||
encoding="utf-8"
|
||||
)
|
||||
hook = (ROOT / "frontend" / "src" / "hooks" / "useDetectionWorkflow.ts").read_text(encoding="utf-8")
|
||||
api = (ROOT / "frontend" / "src" / "services" / "api" / "detection.ts").read_text(encoding="utf-8")
|
||||
types = (ROOT / "frontend" / "src" / "types.ts").read_text(encoding="utf-8")
|
||||
app = (ROOT / "frontend" / "src" / "App.tsx").read_text(encoding="utf-8")
|
||||
provider_panel = (ROOT / "frontend" / "src" / "components" / "providers" / "ProviderPanel.tsx").read_text(
|
||||
encoding="utf-8"
|
||||
)
|
||||
|
||||
assert "interface ModelAssetRead" in types
|
||||
assert "model_asset_id?: string | null" in types
|
||||
assert "listModelAssets" in api
|
||||
assert "/api/v1/detection/model-assets" in api
|
||||
assert "modelAssets" in hook
|
||||
assert "selectedModelAssetId" in hook
|
||||
assert "model_asset_id: selectedModelAssetId || null" in hook
|
||||
assert "Local model assets" in lab
|
||||
assert "onSelectModelAsset" in lab
|
||||
assert "modelAssets={modelAssets}" in app
|
||||
assert "Official reference sources" in provider_panel
|
||||
assert "not AI model choices" in provider_panel
|
||||
|
||||
@@ -66,6 +66,7 @@ def test_configure_yolo_model_dry_run_selects_single_model(tmp_path: Path) -> No
|
||||
assert payload["selected_container_model_path"] == "/app/models/nested/detector.pt"
|
||||
assert payload["env_updates"]["GEOINTEL_INSTALL_AI"] == "true"
|
||||
assert payload["env_updates"]["YOLO_ENABLED"] == "true"
|
||||
assert payload["env_updates"]["YOLO_MODELS_DIR"] == "/app/models"
|
||||
assert payload["env_updates"]["YOLO_MODEL_PATH"] == "/app/models/nested/detector.pt"
|
||||
assert payload["will_download_models"] is False
|
||||
assert not (tmp_path / ".env").exists()
|
||||
@@ -94,4 +95,5 @@ def test_configure_yolo_model_apply_updates_existing_env_file(tmp_path: Path) ->
|
||||
assert "GEOINTEL_FRONTEND_PORT=1202" in contents
|
||||
assert "GEOINTEL_INSTALL_AI=true" in contents
|
||||
assert "YOLO_ENABLED=true" in contents
|
||||
assert "YOLO_MODELS_DIR=/app/models" in contents
|
||||
assert "YOLO_MODEL_PATH=/app/models/detector.engine" in contents
|
||||
|
||||
Reference in New Issue
Block a user