from fastapi import APIRouter, Depends, HTTPException from pydantic import BaseModel from sqlalchemy import select from sqlalchemy.orm import Session from app.core.auth_dependency import require_session from app.db import get_db from app.models.category import Category from app.models.channel import Channel from app.models.channel_category import channel_categories router = APIRouter(dependencies=[Depends(require_session)]) class ChannelCategoriesUpdate(BaseModel): category_ids: list[int] def _category_ids_by_channel(db: Session, channel_ids: list[int]) -> dict[int, list[int]]: if not channel_ids: return {} rows = db.execute( select(channel_categories.c.channel_id, channel_categories.c.category_id).where( channel_categories.c.channel_id.in_(channel_ids) ) ).all() result: dict[int, list[int]] = {} for channel_id, category_id in rows: result.setdefault(channel_id, []).append(category_id) return result def _serialize(channel: Channel, category_ids: list[int]) -> dict: return { "id": channel.id, "youtube_channel_id": channel.youtube_channel_id, "title": channel.title, "description": channel.description, "thumbnail_url": channel.thumbnail_url, "uploads_playlist_id": channel.uploads_playlist_id, "subscribed": channel.subscribed, "last_synced_at": channel.last_synced_at, "category_ids": category_ids, } @router.get("/channels") def list_channels( subscribed: bool | None = None, search: str | None = None, category_id: int | None = None, uncategorized: bool = False, db: Session = Depends(get_db), ) -> list[dict]: query = db.query(Channel) if subscribed is not None: query = query.filter(Channel.subscribed == subscribed) if search: query = query.filter(Channel.title.ilike(f"%{search}%")) if uncategorized: categorized_ids = select(channel_categories.c.channel_id) query = query.filter(~Channel.id.in_(categorized_ids)) elif category_id is not None: channel_ids_in_category = select(channel_categories.c.channel_id).where( channel_categories.c.category_id == category_id ) query = query.filter(Channel.id.in_(channel_ids_in_category)) channels = query.order_by(Channel.title.asc()).all() category_map = _category_ids_by_channel(db, [c.id for c in channels]) return [_serialize(c, category_map.get(c.id, [])) for c in channels] @router.get("/channels/{channel_id}") def get_channel(channel_id: int, db: Session = Depends(get_db)) -> dict: channel = db.get(Channel, channel_id) if channel is None: raise HTTPException(status_code=404, detail="Channel not found") category_map = _category_ids_by_channel(db, [channel_id]) return _serialize(channel, category_map.get(channel_id, [])) @router.put("/channels/{channel_id}/categories") def set_channel_categories(channel_id: int, payload: ChannelCategoriesUpdate, db: Session = Depends(get_db)) -> dict: channel = db.get(Channel, channel_id) if channel is None: raise HTTPException(status_code=404, detail="Channel not found") unique_ids = set(payload.category_ids) if unique_ids: found = db.query(Category.id).filter(Category.id.in_(unique_ids)).all() found_ids = {row[0] for row in found} missing = unique_ids - found_ids if missing: raise HTTPException(status_code=400, detail=f"Unknown category ids: {sorted(missing)}") db.execute(channel_categories.delete().where(channel_categories.c.channel_id == channel_id)) if unique_ids: db.execute( channel_categories.insert(), [{"channel_id": channel_id, "category_id": cid} for cid in unique_ids], ) db.commit() return _serialize(channel, sorted(unique_ids))