156 lines
4.6 KiB
Python
156 lines
4.6 KiB
Python
"""Tests for model registry (src/model_registry.py)."""
|
|
|
|
import pytest
|
|
from src.model_registry import (
|
|
scan_models,
|
|
scan_model_groups,
|
|
ModelConfig,
|
|
ModelGroup,
|
|
_sort_formats,
|
|
FORMAT_PREFERENCE,
|
|
)
|
|
|
|
|
|
def test_scan_returns_list():
|
|
result = scan_models("/nonexistent/path")
|
|
assert isinstance(result, list)
|
|
|
|
|
|
def test_scan_empty_dir(tmp_path):
|
|
result = scan_models(str(tmp_path))
|
|
assert result == []
|
|
|
|
|
|
def test_scan_finds_pt_files(tmp_path):
|
|
(tmp_path / "best.pt").write_bytes(b"fake")
|
|
(tmp_path / "truck-detector.pt").write_bytes(b"fake")
|
|
result = scan_models(str(tmp_path))
|
|
assert len(result) == 2
|
|
names = {m.filename for m in result}
|
|
assert "best.pt" in names
|
|
assert "truck-detector.pt" in names
|
|
|
|
|
|
def test_scan_skips_non_model_files(tmp_path):
|
|
(tmp_path / "modelREADME.md").write_text("readme")
|
|
(tmp_path / "best.pt").write_bytes(b"fake")
|
|
result = scan_models(str(tmp_path))
|
|
assert len(result) == 1
|
|
|
|
|
|
def test_model_config_fields(tmp_path):
|
|
(tmp_path / "v4-best.pt").write_bytes(b"fake")
|
|
result = scan_models(str(tmp_path))
|
|
cfg = result[0]
|
|
assert cfg.filename == "v4-best.pt"
|
|
assert cfg.path == str(tmp_path / "v4-best.pt")
|
|
assert isinstance(cfg.known_classes, list)
|
|
|
|
|
|
def test_model_config_fallback_classes(tmp_path):
|
|
(tmp_path / "unknown-model.pt").write_bytes(b"fake")
|
|
result = scan_models(str(tmp_path))
|
|
cfg = result[0]
|
|
assert cfg.known_classes == []
|
|
|
|
|
|
# --- scan_model_groups tests ---
|
|
|
|
|
|
def test_scan_model_groups_returns_list():
|
|
result = scan_model_groups("/nonexistent/path")
|
|
assert isinstance(result, list)
|
|
|
|
|
|
def test_scan_model_groups_empty_dir(tmp_path):
|
|
result = scan_model_groups(str(tmp_path))
|
|
assert result == []
|
|
|
|
|
|
def test_scan_model_groups_single_stem(tmp_path):
|
|
(tmp_path / "best.pt").write_bytes(b"fake")
|
|
(tmp_path / "best.onnx").write_bytes(b"fake")
|
|
(tmp_path / "best.engine").write_bytes(b"fake")
|
|
result = scan_model_groups(str(tmp_path))
|
|
assert len(result) == 1
|
|
g = result[0]
|
|
assert g.stem == "best"
|
|
assert ".engine" in g.formats
|
|
assert ".pt" in g.formats
|
|
assert ".onnx" in g.formats
|
|
assert g.default_format == ".engine"
|
|
|
|
|
|
def test_scan_model_groups_format_preference(tmp_path):
|
|
(tmp_path / "model.pt").write_bytes(b"fake")
|
|
(tmp_path / "model.onnx").write_bytes(b"fake")
|
|
result = scan_model_groups(str(tmp_path))
|
|
g = result[0]
|
|
assert g.formats == [".pt", ".onnx"]
|
|
assert g.default_format == ".pt"
|
|
|
|
|
|
def test_scan_model_groups_sorted_by_stem(tmp_path):
|
|
(tmp_path / "z-model.pt").write_bytes(b"fake")
|
|
(tmp_path / "a-model.pt").write_bytes(b"fake")
|
|
(tmp_path / "m-model.pt").write_bytes(b"fake")
|
|
result = scan_model_groups(str(tmp_path))
|
|
stems = [g.stem for g in result]
|
|
assert stems == sorted(stems)
|
|
|
|
|
|
def test_scan_model_groups_known_classes(tmp_path):
|
|
(tmp_path / "truck-detector.pt").write_bytes(b"fake")
|
|
(tmp_path / "truck-detector.onnx").write_bytes(b"fake")
|
|
result = scan_model_groups(str(tmp_path))
|
|
g = result[0]
|
|
assert g.known_classes == ["truck"]
|
|
|
|
|
|
def test_scan_model_groups_unknown_stem(tmp_path):
|
|
(tmp_path / "unknown-model.pt").write_bytes(b"fake")
|
|
result = scan_model_groups(str(tmp_path))
|
|
g = result[0]
|
|
assert g.known_classes == []
|
|
|
|
|
|
def test_scan_model_groups_format_paths(tmp_path):
|
|
(tmp_path / "best.pt").write_bytes(b"fake")
|
|
(tmp_path / "best.engine").write_bytes(b"fake")
|
|
result = scan_model_groups(str(tmp_path))
|
|
g = result[0]
|
|
assert ".pt" in g.format_paths
|
|
assert ".engine" in g.format_paths
|
|
assert g.format_paths[".pt"] == str((tmp_path / "best.pt").resolve())
|
|
assert g.format_paths[".engine"] == str((tmp_path / "best.engine").resolve())
|
|
|
|
|
|
def test_scan_model_groups_skips_non_model_files(tmp_path):
|
|
(tmp_path / "best.pt").write_bytes(b"fake")
|
|
(tmp_path / "readme.md").write_text("hi")
|
|
result = scan_model_groups(str(tmp_path))
|
|
assert len(result) == 1
|
|
assert result[0].stem == "best"
|
|
|
|
|
|
def test_scan_model_groups_multiple_stems(tmp_path):
|
|
(tmp_path / "best.pt").write_bytes(b"fake")
|
|
(tmp_path / "v4-best.engine").write_bytes(b"fake")
|
|
(tmp_path / "truck-detector.onnx").write_bytes(b"fake")
|
|
result = scan_model_groups(str(tmp_path))
|
|
assert len(result) == 3
|
|
stems = [g.stem for g in result]
|
|
assert "best" in stems
|
|
assert "v4-best" in stems
|
|
assert "truck-detector" in stems
|
|
|
|
|
|
def test_sort_formats_explicit():
|
|
result = _sort_formats([".onnx", ".pt", ".engine"])
|
|
assert result == [".engine", ".pt", ".onnx"]
|
|
|
|
|
|
def test_sort_formats_unknown_ext():
|
|
result = _sort_formats([".pt", ".xyz", ".onnx"])
|
|
assert result == [".pt", ".onnx", ".xyz"]
|