| """Tests for the batch loader service (spec §4).""" |
| from pathlib import Path |
|
|
| import pytest |
| from sqlalchemy import select |
|
|
| from app.db import SessionLocal |
| from app.models import Case, Dataset, Embedding, KnownDog, Picture, UnknownDog, User |
| from app.models.base import CaseType, KnownDogStatus, SubjectType |
| from app.services.batch_loader import load_dataset |
| from scripts.make_sample_images import make_image |
| from scripts.prepare_test_data import prepare |
|
|
|
|
| @pytest.fixture |
| def db(): |
| session = SessionLocal() |
| try: |
| yield session |
| finally: |
| session.rollback() |
| session.close() |
|
|
|
|
| def _build_input(tmp: Path, layout: dict[str, int], *, identical: bool = False) -> Path: |
| """Write dog folders. identical=True writes the same bytes to every image in a folder so the |
| mock embedder produces matching vectors across the known/found split.""" |
| root = tmp / "input" |
| for folder, count in layout.items(): |
| d = root / folder |
| d.mkdir(parents=True) |
| shared = make_image(abs(hash(folder)) % 1000) |
| for i in range(count): |
| (d / f"img{i}.jpg").write_bytes(shared if identical else make_image((abs(hash(folder)) + i) % 1000)) |
| return root |
|
|
|
|
| def test_load_known_dataset(tmp_path, db): |
| root = _build_input(tmp_path, {"dogA": 2, "dogB": 2}) |
| prepare(root, root, holdout=1, seed=1) |
| result = load_dataset( |
| db, folder=root, dataset_type="known", name="K", description="d", |
| csv_path=root / "known_dogs.csv", |
| ) |
| |
| assert result.dogs_loaded == 2 |
| assert result.images_processed == 2 |
| assert not result.errors |
|
|
| dataset = db.get(Dataset, result.dataset_id) |
| assert dataset.dog_count == 2 |
| dogs = db.execute(select(KnownDog).where(KnownDog.dataset_id == dataset.id)).scalars().all() |
| assert len(dogs) == 2 |
| assert all(d.status == KnownDogStatus.home for d in dogs) |
| |
| owners = db.execute(select(User).where(User.dataset_id == dataset.id)).scalars().all() |
| assert len(owners) == 2 |
| |
| pics = db.execute( |
| select(Picture).where(Picture.subject_type == SubjectType.known, |
| Picture.subject_id.in_([d.id for d in dogs])) |
| ).scalars().all() |
| assert len(pics) == 2 |
| embs = db.execute(select(Embedding).where(Embedding.picture_id.in_([p.id for p in pics]))).scalars().all() |
| assert len(embs) == 2 |
|
|
|
|
| def test_folder_images_grouped_into_one_dog(tmp_path, db): |
| """A folder with several registration images -> ONE KnownDog with multiple pictures.""" |
| root = _build_input(tmp_path, {"dogA": 5}) |
| prepare(root, root, holdout=None, seed=1) |
| result = load_dataset( |
| db, folder=root, dataset_type="known", name="Grouped", description=None, |
| csv_path=root / "known_dogs.csv", |
| ) |
| assert result.dogs_loaded == 1 |
| dogs = db.execute(select(KnownDog).where(KnownDog.dataset_id == result.dataset_id)).scalars().all() |
| assert len(dogs) == 1 |
| pics = db.execute( |
| select(Picture).where(Picture.subject_type == SubjectType.known, Picture.subject_id == dogs[0].id) |
| ).scalars().all() |
| assert len(pics) == 3 |
| assert sum(1 for p in pics if p.is_primary) == 1 |
|
|
|
|
| def test_one_dog_per_image_legacy_mode(tmp_path, db): |
| root = _build_input(tmp_path, {"dogA": 5}) |
| prepare(root, root, holdout=None, seed=1) |
| result = load_dataset( |
| db, folder=root, dataset_type="known", name="Legacy", description=None, |
| csv_path=root / "known_dogs.csv", group_by_folder=False, |
| ) |
| assert result.dogs_loaded == 3 |
|
|
|
|
| def test_load_unknown_dataset_found_zip_matches(tmp_path, db): |
| root = _build_input(tmp_path, {"dogA": 3}) |
| prepare(root, root, holdout=1, seed=2) |
| |
| import csv as _csv |
| known_zip = next(_csv.DictReader((root / "known_dogs.csv").open()))["zip"] |
|
|
| result = load_dataset( |
| db, folder=root, dataset_type="unknown", name="U", description="d", |
| csv_path=root / "found_dogs.csv", |
| ) |
| assert result.dogs_loaded == 1 |
| assert result.cases_created == 1 |
| unknown = db.execute(select(UnknownDog).where(UnknownDog.dataset_id == result.dataset_id)).scalars().one() |
| case = db.execute(select(Case).where(Case.unknown_dog_id == unknown.id)).scalars().one() |
| assert case.type == CaseType.found |
| assert case.event_zip == known_zip |
|
|
|
|
| def test_mark_lost_flag(tmp_path, db): |
| root = _build_input(tmp_path, {"dogA": 2, "dogB": 2}) |
| prepare(root, root, holdout=1, seed=1) |
| result = load_dataset( |
| db, folder=root, dataset_type="known", name="K", description="d", |
| csv_path=root / "known_dogs.csv", mark_lost=True, mark_lost_pct=100, |
| ) |
| dogs = db.execute(select(KnownDog).where(KnownDog.dataset_id == result.dataset_id)).scalars().all() |
| assert all(d.status == KnownDogStatus.lost for d in dogs) |
| |
| assert result.cases_created == len(dogs) |
| lost_cases = db.execute( |
| select(Case).where(Case.type == CaseType.lost, Case.known_dog_id.in_([d.id for d in dogs])) |
| ).scalars().all() |
| assert len(lost_cases) == len(dogs) |
|
|
|
|
| def test_bad_image_rows_skipped(tmp_path, db): |
| root = _build_input(tmp_path, {"dogA": 2}) |
| prepare(root, root, holdout=1, seed=1) |
| |
| csv_path = root / "known_dogs.csv" |
| text = csv_path.read_text().rstrip("\n") |
| csv_path.write_text(text + "\nnope_folder,missing.jpg,Ghost,black,large,x,20001,N,n@example.com,555\n") |
|
|
| result = load_dataset( |
| db, folder=root, dataset_type="known", name="K", csv_path=csv_path, description=None, |
| ) |
| assert result.dogs_loaded == 1 |
| assert len(result.errors) == 1 |
| assert "not found" in result.errors[0] |
|
|
|
|
| def test_run_matching_after_load(tmp_path, db): |
| |
| root = _build_input(tmp_path, {"dogA": 2}, identical=True) |
| prepare(root, root, holdout=1, seed=1) |
| load_dataset(db, folder=root, dataset_type="known", name="K", |
| csv_path=root / "known_dogs.csv", description=None, mark_lost=True, mark_lost_pct=100) |
| found = load_dataset(db, folder=root, dataset_type="unknown", name="U", |
| csv_path=root / "found_dogs.csv", description=None, run_matching=True) |
| assert found.matching is not None |
| assert found.matching["matches_created"] >= 1 |
|
|