from fastapi import APIRouter, Depends, HTTPException from pydantic import BaseModel, Field from sqlalchemy import func from sqlalchemy.exc import IntegrityError from sqlalchemy.orm import Session from app.core.auth_dependency import require_session from app.core.slugify import unique_slugify from app.db import get_db from app.models.category import Category from app.models.channel_category import channel_categories router = APIRouter(dependencies=[Depends(require_session)]) class CategoryCreate(BaseModel): name: str = Field(min_length=1, max_length=255) class CategoryUpdate(BaseModel): name: str = Field(min_length=1, max_length=255) class CategoryReorder(BaseModel): category_ids: list[int] def _existing_slugs(db: Session, exclude_id: int | None = None) -> set[str]: query = db.query(Category.slug) if exclude_id is not None: query = query.filter(Category.id != exclude_id) return {row[0] for row in query.all()} def _name_taken(db: Session, name: str, exclude_id: int | None = None) -> bool: query = db.query(Category.name) if exclude_id is not None: query = query.filter(Category.id != exclude_id) target = name.casefold() return any(row[0].casefold() == target for row in query.all()) def _serialize(db: Session, category: Category, counts: dict[int, int]) -> dict: return { "id": category.id, "name": category.name, "slug": category.slug, "sort_order": category.sort_order, "channel_count": counts.get(category.id, 0), } @router.get("/categories") def list_categories(db: Session = Depends(get_db)) -> list[dict]: categories = db.query(Category).order_by(Category.sort_order.asc(), Category.id.asc()).all() count_rows = ( db.query(channel_categories.c.category_id, func.count(channel_categories.c.channel_id)) .group_by(channel_categories.c.category_id) .all() ) counts = dict(count_rows) return [_serialize(db, c, counts) for c in categories] @router.post("/categories", status_code=201) def create_category(payload: CategoryCreate, db: Session = Depends(get_db)) -> dict: name = payload.name.strip() if not name: raise HTTPException(status_code=400, detail="Category name must not be empty") if _name_taken(db, name): raise HTTPException(status_code=409, detail="Category with this name already exists") slug = unique_slugify(name, _existing_slugs(db)) max_sort_order = db.query(func.max(Category.sort_order)).scalar() or 0 category = Category(name=name, slug=slug, sort_order=max_sort_order + 1) db.add(category) try: db.commit() except IntegrityError: db.rollback() raise HTTPException(status_code=409, detail="Category with this name already exists") return _serialize(db, category, {}) @router.patch("/categories/{category_id}") def update_category(category_id: int, payload: CategoryUpdate, db: Session = Depends(get_db)) -> dict: category = db.get(Category, category_id) if category is None: raise HTTPException(status_code=404, detail="Category not found") name = payload.name.strip() if not name: raise HTTPException(status_code=400, detail="Category name must not be empty") if _name_taken(db, name, exclude_id=category_id): raise HTTPException(status_code=409, detail="Category with this name already exists") category.name = name category.slug = unique_slugify(name, _existing_slugs(db, exclude_id=category_id)) try: db.commit() except IntegrityError: db.rollback() raise HTTPException(status_code=409, detail="Category with this name already exists") return _serialize(db, category, {}) @router.delete("/categories/{category_id}", status_code=204) def delete_category(category_id: int, db: Session = Depends(get_db)) -> None: category = db.get(Category, category_id) if category is None: raise HTTPException(status_code=404, detail="Category not found") db.delete(category) db.commit() @router.post("/categories/reorder") def reorder_categories(payload: CategoryReorder, db: Session = Depends(get_db)) -> list[dict]: categories = {c.id: c for c in db.query(Category).all()} if set(payload.category_ids) != set(categories.keys()): raise HTTPException(status_code=400, detail="category_ids must contain exactly all existing category ids") for index, category_id in enumerate(payload.category_ids): categories[category_id].sort_order = index db.commit() ordered = db.query(Category).order_by(Category.sort_order.asc(), Category.id.asc()).all() return [_serialize(db, c, {}) for c in ordered]