"""Resolve channel catalog for a given user based on their assigned categories.""" from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import selectinload from ..models.user import User, UserCatalogEntry, UserProviderPreference from ..models.channel import Channel, Category, ChannelProviderMap, ContentType from ..models.provider import ProviderAccount async def get_user_categories(db: AsyncSession, user: User, content_type: ContentType | None = None) -> list[Category]: q = ( select(Category) .join(UserCatalogEntry, UserCatalogEntry.category_id == Category.id) .where(UserCatalogEntry.user_id == user.id) ) if content_type: q = q.where(Category.type == content_type) result = await db.execute(q) return list(result.scalars().all()) async def get_user_channels( db: AsyncSession, user: User, content_type: ContentType | None = None, category_id: int | None = None, ) -> list[Channel]: category_ids_q = ( select(UserCatalogEntry.category_id) .where(UserCatalogEntry.user_id == user.id) ) q = ( select(Channel) .options(selectinload(Channel.provider_maps), selectinload(Channel.category)) .where(Channel.category_id.in_(category_ids_q)) .where(Channel.is_active == True) # noqa: E712 ) if content_type: q = q.where(Channel.type == content_type) if category_id: q = q.where(Channel.category_id == category_id) result = await db.execute(q) return list(result.scalars().all()) async def get_channel_provider_maps(db: AsyncSession, channel_id: int) -> list[dict]: q = ( select(ChannelProviderMap, ProviderAccount.name.label("provider_name")) .join(ProviderAccount, ProviderAccount.id == ChannelProviderMap.provider_account_id) .where(ChannelProviderMap.channel_id == channel_id) ) result = await db.execute(q) rows = result.all() return [ { "provider_account_id": row.ChannelProviderMap.provider_account_id, "stream_url": row.ChannelProviderMap.stream_url, "provider_name": row.provider_name or f"#{row.ChannelProviderMap.provider_account_id}", } for row in rows ] async def can_user_access_channel(db: AsyncSession, user: User, channel_id: int) -> bool: channels = await get_user_channels(db, user) return any(ch.id == channel_id for ch in channels) async def get_user_preferred_provider_ids(db: AsyncSession, user: User) -> list[int]: result = await db.execute( select(UserProviderPreference.provider_account_id) .where(UserProviderPreference.user_id == user.id) .order_by(UserProviderPreference.priority) ) return list(result.scalars().all())