from __future__ import annotations from collections.abc import Mapping, Sequence from dataclasses import dataclass from typing import Any import pytest from cmvr_edge_ai.contracts import Detection, ImageFrame from cmvr_edge_ai.detection.base import DetectionModel from cmvr_edge_ai.detection.registry import ( DetectionModelRegistry, DetectionModelRegistryError, DetectionModelSpec, DuplicateDetectionModelError, UnknownDetectionModelError, ) class FakeDetectionModel(DetectionModel): def load(self) -> None: pass def predict( self, frame: ImageFrame, labels: tuple[str, ...], confidence: float, ) -> Sequence[Detection]: del frame, labels, confidence return () def close(self) -> None: pass def _factory(params: Mapping[str, Any]) -> DetectionModel: del params return FakeDetectionModel() def _spec(model_id: str = "ppe-yolo@1") -> DetectionModelSpec: return DetectionModelSpec( model_id=model_id, name="Construction PPE", supported_labels=("Worker", "No-Helmet", "No-Vest"), factory=_factory, backend="ultralytics", description="PPE violation detector", ) def test_detection_model_spec_preserves_order_and_copies_labels() -> None: labels = ["Worker", "No-Helmet", "No-Vest"] spec = DetectionModelSpec( model_id="ppe@2", name=" PPE detector ", supported_labels=labels, # type: ignore[arg-type] factory=_factory, backend=" yolo ", ) labels.append("No-Glove") assert spec.name == "PPE detector" assert spec.backend == "yolo" assert spec.supported_labels == ("Worker", "No-Helmet", "No-Vest") @pytest.mark.parametrize("model_id", ["ppe", "@1", "ppe@0", "ppe@-1", "ppe@v1"]) def test_detection_model_spec_requires_positive_version(model_id: str) -> None: with pytest.raises(ValueError, match="model id"): DetectionModelSpec( model_id=model_id, name="PPE", supported_labels=("Worker",), factory=_factory, backend="fake", ) @pytest.mark.parametrize( ("overrides", "exception", "message"), [ ({"name": " "}, ValueError, "model name"), ({"backend": ""}, ValueError, "model backend"), ({"supported_labels": ()}, ValueError, "must not be empty"), ({"supported_labels": ("Worker", "")}, ValueError, "non-empty strings"), ( {"supported_labels": ("Worker", "Worker")}, ValueError, "must be unique", ), ({"supported_labels": "Worker"}, TypeError, "ordered collection"), ({"supported_labels": {"Worker"}}, TypeError, "ordered collection"), ({"factory": None}, TypeError, "factory must be callable"), ], ) def test_detection_model_spec_rejects_invalid_metadata( overrides: dict[str, Any], exception: type[Exception], message: str ) -> None: values: dict[str, Any] = { "model_id": "ppe@1", "name": "PPE", "supported_labels": ("Worker",), "factory": _factory, "backend": "fake", } values.update(overrides) with pytest.raises(exception, match=message): DetectionModelSpec(**values) def test_registry_registers_resolves_and_sorts_specs() -> None: registry = DetectionModelRegistry() second = registry.register(_spec("z-model@1")) first = registry.register(_spec("a-model@2")) assert registry.resolve("z-model@1") is second assert registry.specs() == (first, second) def test_registry_rejects_duplicate_and_reports_available_models() -> None: registry = DetectionModelRegistry() registry.register(_spec("ppe@1")) with pytest.raises(DuplicateDetectionModelError, match="ppe@1"): registry.register(_spec("ppe@1")) with pytest.raises(UnknownDetectionModelError, match=r"available models: ppe@1"): registry.resolve("missing@1") @dataclass class _EntryPoint: name: str value: Any def load(self) -> Any: return self.value class _EntryPoints(list[_EntryPoint]): def select(self, *, group: str) -> _EntryPoints: assert group == DetectionModelRegistry.ENTRY_POINT_GROUP return self def test_registry_discovers_specs_and_registration_callbacks(monkeypatch) -> None: # type: ignore[no-untyped-def] def register_callback(registry: DetectionModelRegistry) -> None: registry.register(_spec("callback@1")) entry_points = _EntryPoints( [ _EntryPoint("spec", _spec("spec@1")), _EntryPoint("callback", register_callback), ] ) monkeypatch.setattr( "cmvr_edge_ai.detection.registry.metadata.entry_points", lambda: entry_points, ) registry = DetectionModelRegistry() registry.load_entry_points() assert [spec.model_id for spec in registry.specs()] == ["callback@1", "spec@1"] def test_registry_rejects_invalid_entry_point(monkeypatch) -> None: # type: ignore[no-untyped-def] entry_points = _EntryPoints([_EntryPoint("invalid", object())]) monkeypatch.setattr( "cmvr_edge_ai.detection.registry.metadata.entry_points", lambda: entry_points, ) with pytest.raises(DetectionModelRegistryError, match="invalid"): DetectionModelRegistry().load_entry_points()