"""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"]