Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
81 changes: 81 additions & 0 deletions app/routers/collections.py
Original file line number Diff line number Diff line change
Expand Up @@ -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),
Expand All @@ -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,
Expand All @@ -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
]
Expand Down Expand Up @@ -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,
Expand All @@ -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, []),
)


Expand Down
8 changes: 8 additions & 0 deletions app/schemas/collection.py
Original file line number Diff line number Diff line change
Expand Up @@ -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] = []
138 changes: 138 additions & 0 deletions tests/test_collections.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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"] == []