"""Admin API: category/catalog management.""" from fastapi import APIRouter, Depends from sqlalchemy import select, func from sqlalchemy.ext.asyncio import AsyncSession from pydantic import BaseModel from ...database import get_db from ...models.channel import Category, Channel, ChannelProviderMap, ContentType from ..auth import get_current_admin router = APIRouter(prefix="/catalogs", tags=["catalogs"]) class CategoryOut(BaseModel): id: int name: str type: ContentType provider_category_id: str | None provider_account_id: int | None class Config: from_attributes = True class GroupedCategoryOut(BaseModel): name: str type: ContentType category_ids: list[int] provider_count: int # distinct providers that can serve channels in this category @router.get("/", response_model=list[CategoryOut]) async def list_categories( content_type: ContentType | None = None, db: AsyncSession = Depends(get_db), _=Depends(get_current_admin), ): q = select(Category).order_by(Category.type, Category.name) if content_type: q = q.where(Category.type == content_type) result = await db.execute(q) return result.scalars().all() @router.get("/grouped", response_model=list[GroupedCategoryOut]) async def list_categories_grouped( content_type: ContentType | None = None, db: AsyncSession = Depends(get_db), _=Depends(get_current_admin), ): """Returns categories merged by name+type. Includes how many providers can serve each.""" q = select(Category).order_by(Category.type, Category.name) if content_type: q = q.where(Category.type == content_type) cats_result = await db.execute(q) all_cats = cats_result.scalars().all() # Group by (name_normalized, type) seen: dict[tuple[str, str], dict] = {} for cat in all_cats: key = (cat.name.strip().lower(), str(cat.type)) if key not in seen: seen[key] = {"name": cat.name, "type": cat.type, "category_ids": [], "provider_count": 0} seen[key]["category_ids"].append(cat.id) if not seen: return [] # Map each category_id back to its group key for the count query cat_id_to_key: dict[int, tuple[str, str]] = {} for key, group in seen.items(): for cat_id in group["category_ids"]: cat_id_to_key[cat_id] = key # Single query: distinct providers per category_id — avoids locale-dependent lower() issues pcount_q = await db.execute( select( Channel.category_id, func.count(func.distinct(ChannelProviderMap.provider_account_id)).label("pcount"), ) .join(ChannelProviderMap, ChannelProviderMap.channel_id == Channel.id) .where(Channel.category_id.in_(list(cat_id_to_key.keys()))) .group_by(Channel.category_id) ) for row in pcount_q: key = cat_id_to_key.get(row.category_id) if key and key in seen: seen[key]["provider_count"] = max(seen[key]["provider_count"], row.pcount) return list(seen.values())