PawTrace / backend /tests /test_skip_embeddings.py
Elliott Duke
HomingPet: lost-dog reunification (FastAPI + React) with Render deploy
de1e3fc
Raw
History Blame Contribute Delete
2.38 kB
"""Loading with skip_embeddings stores images only; embed-all generates embeddings later."""
from pathlib import Path
import pytest
from sqlalchemy import func, select
from app.db import SessionLocal
from app.models import BreedPrediction, Dataset, Embedding, Picture
from app.services.batch_loader import load_dataset
from app.services.datasets import embed_all
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) -> Path:
root = tmp / "input"
for folder in ("dogA", "dogB"):
d = root / folder
d.mkdir(parents=True)
for i in range(2):
(d / f"img{i}.jpg").write_bytes(make_image(abs(hash((folder, i))) % 1000))
return root
def test_load_skips_embeddings_then_embed_all_generates(tmp_path, db):
root = _build_input(tmp_path)
prepare(root, root, holdout=1, seed=1)
result = load_dataset(
db, folder=root, dataset_type="known", name="StorageOnly", description=None,
csv_path=root / "known_dogs.csv", skip_embeddings=True,
)
dataset = db.get(Dataset, result.dataset_id)
# Pictures were stored, but no embeddings or breed predictions were generated.
pic_ids = db.execute(
select(Picture.id).where(Picture.subject_type == "known")
).scalars().all()
assert len(pic_ids) == result.images_processed >= 1
assert db.execute(select(func.count()).select_from(Embedding)).scalar_one() == 0
assert db.execute(select(func.count()).select_from(BreedPrediction)).scalar_one() == 0
# embed-all backfills embeddings for the stored images.
summary = embed_all(db, dataset)
assert summary["embedded"] == summary["total_pictures"] == len(pic_ids)
assert db.execute(select(func.count()).select_from(Embedding)).scalar_one() == len(pic_ids)
def test_default_load_still_embeds(tmp_path, db):
root = _build_input(tmp_path)
prepare(root, root, holdout=1, seed=1)
load_dataset(
db, folder=root, dataset_type="known", name="WithEmb", description=None,
csv_path=root / "known_dogs.csv", # skip_embeddings defaults False
)
assert db.execute(select(func.count()).select_from(Embedding)).scalar_one() >= 1