diff --git a/app/routers/collections.py b/app/routers/collections.py index ae84820..3d65820 100644 --- a/app/routers/collections.py +++ b/app/routers/collections.py @@ -7,16 +7,90 @@ from app.auth import current_active_user from app.database import get_async_session from app.models import Collection, Item, User +from app.models.image import Image from app.schemas.collection import ( CollectionCreate, CollectionRead, CollectionReadWithCount, CollectionUpdate, + ImagePreview, ) router = APIRouter(prefix="/collections", tags=["collections"]) +async def _get_preview_images( + collection_ids: list[str], + session: AsyncSession, +) -> dict[str, list[ImagePreview]]: + """Fetch up to 4 preview images per collection. + + For each collection, picks the first image (by position) from each of the + first 4 items that have images, ordered by item creation date. + """ + if not collection_ids: + return {} + + # Use a window function to rank images within each collection, + # picking the first image per item (by position), then ranking items by created_at. + from sqlalchemy import and_ + + # Subquery: for each image, get its collection_id via the item, + # and rank it: first by item created_at, then by image position. + # We only want the first image per item, then the first 4 items per collection. + first_image_per_item = ( + select( + Image.id.label("image_id"), + Image.url.label("image_url"), + Item.collection_id.label("collection_id"), + func.row_number() + .over( + partition_by=[Item.collection_id, Item.id], + order_by=Image.position, + ) + .label("img_rank"), + func.min(Image.position) + .over(partition_by=[Item.collection_id, Item.id]) + .label("_min_pos"), + ) + .join(Item, and_(Image.item_id == Item.id, Item.collection_id.in_(collection_ids))) + .subquery() + ) + + # Filter to first image per item, then rank items within each collection + ranked = ( + select( + first_image_per_item.c.image_id, + first_image_per_item.c.image_url, + first_image_per_item.c.collection_id, + func.row_number() + .over( + partition_by=first_image_per_item.c.collection_id, + order_by=first_image_per_item.c.image_id, + ) + .label("item_rank"), + ) + .where(first_image_per_item.c.img_rank == 1) + .subquery() + ) + + stmt = ( + select(ranked.c.collection_id, ranked.c.image_id, ranked.c.image_url) + .where(ranked.c.item_rank <= 4) + .order_by(ranked.c.collection_id, ranked.c.item_rank) + ) + + result = await session.execute(stmt) + rows = result.all() + + previews: dict[str, list[ImagePreview]] = {} + for row in rows: + previews.setdefault(row.collection_id, []).append( + ImagePreview(id=row.image_id, url=row.image_url) + ) + return previews + + @router.get("", response_model=list[CollectionReadWithCount]) async def list_collections( user: User = Depends(current_active_user), @@ -34,6 +108,9 @@ async def list_collections( result = await session.execute(stmt) rows = result.all() + collection_ids = [row.Collection.id for row in rows] + previews = await _get_preview_images(collection_ids, session) + return [ CollectionReadWithCount( id=row.Collection.id, @@ -44,6 +121,7 @@ async def list_collections( created_at=row.Collection.created_at, updated_at=row.Collection.updated_at, item_count=row.item_count, + preview_images=previews.get(row.Collection.id, []), ) for row in rows ] @@ -71,6 +149,8 @@ async def get_collection( detail="Collection not found", ) + previews = await _get_preview_images([row.Collection.id], session) + return CollectionReadWithCount( id=row.Collection.id, user_id=row.Collection.user_id, @@ -80,6 +160,7 @@ async def get_collection( created_at=row.Collection.created_at, updated_at=row.Collection.updated_at, item_count=row.item_count, + preview_images=previews.get(row.Collection.id, []), ) diff --git a/app/schemas/collection.py b/app/schemas/collection.py index 9d596bb..1e3155f 100644 --- a/app/schemas/collection.py +++ b/app/schemas/collection.py @@ -53,7 +53,15 @@ class CollectionRead(CollectionBase): updated_at: datetime +class ImagePreview(BaseModel): + """Lightweight image preview for collection cards.""" + + id: UUID + url: str + + class CollectionReadWithCount(CollectionRead): """Schema for reading a Collection with item count.""" item_count: int = 0 + preview_images: list[ImagePreview] = [] diff --git a/tests/test_collections.py b/tests/test_collections.py index 8af545b..ee539ab 100644 --- a/tests/test_collections.py +++ b/tests/test_collections.py @@ -3,6 +3,7 @@ from sqlalchemy.ext.asyncio import AsyncSession from app.models import Collection, Item, User +from app.models.image import Image @pytest.mark.asyncio @@ -289,3 +290,140 @@ async def test_collections_isolation( data = response.json() assert len(data) == 1 assert data[0]["name"] == "My Collection" + + +@pytest.mark.asyncio +async def test_list_collections_preview_images( + client: AsyncClient, session: AsyncSession, test_user: User, auth_client +): + """Test that listing collections includes preview images from items.""" + collection = Collection(user_id=str(test_user.id), name="My Collection") + session.add(collection) + await session.commit() + await session.refresh(collection) + + # Create items with images + item1 = Item(user_id=str(test_user.id), collection_id=collection.id, name="Item 1") + item2 = Item(user_id=str(test_user.id), collection_id=collection.id, name="Item 2") + session.add_all([item1, item2]) + await session.commit() + await session.refresh(item1) + await session.refresh(item2) + + img1 = Image( + user_id=str(test_user.id), + item_id=item1.id, + filename="a.jpg", + storage_key="key-a", + url="https://example.com/a.jpg", + content_type="image/jpeg", + size_bytes=1024, + position=0, + ) + img2 = Image( + user_id=str(test_user.id), + item_id=item2.id, + filename="b.jpg", + storage_key="key-b", + url="https://example.com/b.jpg", + content_type="image/jpeg", + size_bytes=1024, + position=0, + ) + session.add_all([img1, img2]) + await session.commit() + + response = await client.get("/collections") + assert response.status_code == 200 + data = response.json() + assert len(data) == 1 + assert len(data[0]["preview_images"]) == 2 + urls = {p["url"] for p in data[0]["preview_images"]} + assert "https://example.com/a.jpg" in urls + assert "https://example.com/b.jpg" in urls + + +@pytest.mark.asyncio +async def test_list_collections_preview_images_max_four( + client: AsyncClient, session: AsyncSession, test_user: User, auth_client +): + """Test that preview images are limited to 4 per collection.""" + collection = Collection(user_id=str(test_user.id), name="Big Collection") + session.add(collection) + await session.commit() + await session.refresh(collection) + + # Create 6 items each with an image + for i in range(6): + item = Item(user_id=str(test_user.id), collection_id=collection.id, name=f"Item {i}") + session.add(item) + await session.commit() + await session.refresh(item) + img = Image( + user_id=str(test_user.id), + item_id=item.id, + filename=f"img{i}.jpg", + storage_key=f"key-{i}", + url=f"https://example.com/img{i}.jpg", + content_type="image/jpeg", + size_bytes=1024, + position=0, + ) + session.add(img) + await session.commit() + + response = await client.get("/collections") + assert response.status_code == 200 + data = response.json() + assert len(data[0]["preview_images"]) == 4 + + +@pytest.mark.asyncio +async def test_get_collection_preview_images( + client: AsyncClient, session: AsyncSession, test_user: User, auth_client +): + """Test that getting a single collection includes preview images.""" + collection = Collection(user_id=str(test_user.id), name="My Collection") + session.add(collection) + await session.commit() + await session.refresh(collection) + + item = Item(user_id=str(test_user.id), collection_id=collection.id, name="Item 1") + session.add(item) + await session.commit() + await session.refresh(item) + + img = Image( + user_id=str(test_user.id), + item_id=item.id, + filename="pic.jpg", + storage_key="key-pic", + url="https://example.com/pic.jpg", + content_type="image/jpeg", + size_bytes=1024, + position=0, + ) + session.add(img) + await session.commit() + + response = await client.get(f"/collections/{collection.id}") + assert response.status_code == 200 + data = response.json() + assert len(data["preview_images"]) == 1 + assert data["preview_images"][0]["url"] == "https://example.com/pic.jpg" + + +@pytest.mark.asyncio +async def test_collection_no_images_empty_preview( + client: AsyncClient, session: AsyncSession, test_user: User, auth_client +): + """Test that collections without images have empty preview_images.""" + collection = Collection(user_id=str(test_user.id), name="Empty Collection") + session.add(collection) + await session.commit() + await session.refresh(collection) + + response = await client.get(f"/collections/{collection.id}") + assert response.status_code == 200 + data = response.json() + assert data["preview_images"] == []