107 lines
3.8 KiB
Python
107 lines
3.8 KiB
Python
|
|
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))
|