cmvr_edge_ai/tests/unit/test_detection_registry.py

178 lines
5.3 KiB
Python
Raw Normal View History

2026-07-20 16:59:37 +08:00
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()