myYouTube/backend/app/api/categories.py

135 lines
4.6 KiB
Python
Raw Normal View History

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]