PawTrace / backend /tests /test_batch_loader.py
Elliott Duke
HomingPet: lost-dog reunification (FastAPI + React) with Render deploy
de1e3fc
Raw
History Blame Contribute Delete
6.87 kB
"""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",
)
# 2 dogs x (2-1 holdout) = 2 known images -> 2 KnownDog rows.
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 created and linked to the dataset.
owners = db.execute(select(User).where(User.dataset_id == dataset.id)).scalars().all()
assert len(owners) == 2
# Image pipeline ran: pictures + embeddings exist.
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}) # 5 imgs -> holdout 2 -> 3 registration images
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 # one identity, not three
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 # all registration images attached to the single dog
assert sum(1 for p in pics if p.is_primary) == 1 # exactly one primary
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 # one dog per registration image (legacy)
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)
# known zip from the CSV
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 # found_zip == known zip -> radius matching works
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)
# Marking lost opens a lost case per dog.
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)
# Corrupt the CSV: append a row pointing at a missing image; load must continue.
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 # the one valid row
assert len(result.errors) == 1 # the missing image row
assert "not found" in result.errors[0]
def test_run_matching_after_load(tmp_path, db):
# Identical images so the mock embedder yields a strong known<->found match at radius 0.
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