178 lines
5.3 KiB
Python
178 lines
5.3 KiB
Python
|
|
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()
|