"""ProviderPool: manages provider account slots and live stream allocation.""" import asyncio import json import logging import time from typing import AsyncIterator import redis.asyncio as aioredis from ..config import settings from .restream import BroadcastGroup, ClientHandle logger = logging.getLogger(__name__) ACTIVE_STREAM_KEY = "ks:stream:{channel_id}" PROVIDER_SLOT_KEY = "ks:slot:{account_id}" USER_CONN_KEY = "ks:user_conns:{user_id}" MONITORING_KEY = "ks:monitoring" class NoSlotsAvailableError(Exception): pass class MaxConnectionsExceededError(Exception): pass class ProviderPool: """ Singleton managing all active BroadcastGroups and provider slot allocation. Uses Redis for cross-process state and in-memory dict for the asyncio objects. """ def __init__(self): self._groups: dict[int, BroadcastGroup] = {} self._lock = asyncio.Lock() self._redis: aioredis.Redis | None = None async def init(self) -> None: self._redis = aioredis.from_url(settings.redis_url, decode_responses=True) await self._redis.ping() # Clear any stale keys left from a previous process crash or restart stale = await self._redis.keys("ks:*") if stale: await self._redis.delete(*stale) logger.info(f"ProviderPool cleared {len(stale)} stale Redis keys on startup") logger.info("ProviderPool initialized, Redis connected") async def close(self) -> None: if self._redis: await self._redis.aclose() # ------------------------------------------------------------------------- # Public interface # ------------------------------------------------------------------------- async def acquire( self, channel_id: int, stream_url_resolver, # async callable(provider_account_id) -> str provider_maps: list[dict], # [{"provider_account_id": X, "stream_url": Y, "provider_name": Z}, ...] user_id: str, user_max_conns: int, is_priority: bool = False, preferred_provider_ids: list[int] | None = None, username: str = "", channel_name: str = "", ) -> ClientHandle: await self._check_user_connections(user_id, user_max_conns) async with self._lock: # Case 1: stream already active → attach to existing group if channel_id in self._groups: group = self._groups[channel_id] handle = await group.add_client(user_id, username=username, is_priority=is_priority) await self._redis.hincrby(ACTIVE_STREAM_KEY.format(channel_id=channel_id), "clients", 1) await self._redis.sadd(USER_CONN_KEY.format(user_id=user_id), handle.client_id) await self._publish_monitoring() logger.info(f"User {user_id} attached to existing stream ch={channel_id}") return handle # Case 2: find a free slot (preferred providers sorted first) sorted_maps = self._sort_by_preference(provider_maps, preferred_provider_ids) account_id, stream_url = await self._find_free_slot(sorted_maps) if account_id is None: if not is_priority: raise NoSlotsAvailableError(f"No free provider slots for channel {channel_id}") # Case 3: priority eviction — bump a non-priority stream account_id, stream_url, victim_group = await self._find_eviction_candidate(sorted_maps) if account_id is None: raise NoSlotsAvailableError(f"No evictable non-priority streams for channel {channel_id}") # Clean Redis for the victim stream inside the lock victim_channel_id = victim_group.channel_id victim_account_id = victim_group.provider_account_id self._groups.pop(victim_channel_id, None) await self._redis.delete(ACTIVE_STREAM_KEY.format(channel_id=victim_channel_id)) await self._redis.delete(PROVIDER_SLOT_KEY.format(account_id=victim_account_id)) # Stop victim group outside lock (disconnect all its clients via SENTINEL) asyncio.create_task(victim_group.stop()) logger.info( f"Priority eviction: stopped ch={victim_channel_id} " f"(provider_account={victim_account_id}) for priority user {user_id}" ) # Reserve slot await self._redis.set(PROVIDER_SLOT_KEY.format(account_id=account_id), channel_id) provider_name = next( (pm.get("provider_name", "") for pm in provider_maps if pm["provider_account_id"] == account_id), "" ) stream_urls = await self._build_stream_urls(account_id, stream_url) group = BroadcastGroup( channel_id=channel_id, stream_urls=stream_urls, provider_account_id=account_id, channel_name=channel_name, provider_name=provider_name, ) self._groups[channel_id] = group handle = await group.add_client(user_id, username=username, is_priority=is_priority) await self._redis.hset( ACTIVE_STREAM_KEY.format(channel_id=channel_id), mapping={ "provider_account_id": account_id, "stream_url": stream_url, "clients": 1, "started_at": time.time(), }, ) await self._redis.sadd(USER_CONN_KEY.format(user_id=user_id), handle.client_id) await group.start() await self._publish_monitoring() logger.info(f"New stream started ch={channel_id} provider_account={account_id}") return handle async def release(self, channel_id: int, client_id: str, user_id: str) -> None: group_to_stop: BroadcastGroup | None = None stopped_account_id: int | None = None async with self._lock: group = self._groups.get(channel_id) if not group: return remaining = await group.remove_client(client_id) await self._redis.srem(USER_CONN_KEY.format(user_id=user_id), client_id) if remaining == 0: # Remove from registry and clean Redis *while holding the lock*, # but defer group.stop() (which awaits task cancellation) until # after the lock is released to avoid blocking the pool. self._groups.pop(channel_id, None) group_to_stop = group stopped_account_id = group.provider_account_id await self._redis.delete(ACTIVE_STREAM_KEY.format(channel_id=channel_id)) await self._redis.delete(PROVIDER_SLOT_KEY.format(account_id=stopped_account_id)) else: await self._redis.hset( ACTIVE_STREAM_KEY.format(channel_id=channel_id), "clients", remaining ) await self._publish_monitoring() if group_to_stop is not None: await group_to_stop.stop() logger.info(f"Stream stopped ch={channel_id}, slot freed: provider_account={stopped_account_id}") async def kill_client(self, client_id: str) -> bool: """Admin: force-disconnect a specific client by sending SENTINEL to their queue.""" async with self._lock: for group in self._groups.values(): if client_id in group._clients: return await group.force_disconnect_client(client_id) return False def get_active_streams(self) -> list[dict]: return [g.stats() for g in self._groups.values()] async def get_user_connection_count(self, user_id: str) -> int: return await self._redis.scard(USER_CONN_KEY.format(user_id=user_id)) # ------------------------------------------------------------------------- # Internal helpers # ------------------------------------------------------------------------- async def _check_user_connections(self, user_id: str, max_conns: int) -> None: # Cross-check Redis set against in-memory groups to auto-heal stale entries redis_conns = await self._redis.smembers(USER_CONN_KEY.format(user_id=user_id)) if redis_conns: active = {cid for g in self._groups.values() for cid in g.client_ids} stale = redis_conns - active if stale: await self._redis.srem(USER_CONN_KEY.format(user_id=user_id), *stale) logger.warning(f"Auto-cleaned {len(stale)} stale conn(s) for user {user_id}") redis_conns -= stale count = len(redis_conns) if count >= max_conns: # Determine if this is a zapping scenario (same device switching channels quickly) # or a genuine multi-device connection attempt. # Zapping: the existing connection is very recent (< 10 s) — the TV app opens the # new stream before the HTTP response of the old one closes. In that case we evict # silently so the channel switch is seamless. # Multi-device: the existing connection is older — keep it alive and raise so the # caller can show a "blocked" stream on the new device instead. now = time.time() groups_snapshot = list(self._groups.values()) oldest_age = 0.0 for group in groups_snapshot: for cid, handle in group._clients.items(): if cid in redis_conns: oldest_age = max(oldest_age, now - handle.connected_at) if oldest_age < 10.0: # Zapping — evict silently logger.info(f"User {user_id} zapping ({oldest_age:.1f}s old conn) — evicting old stream") for client_id in list(redis_conns): for group in groups_snapshot: if client_id in group._clients: await group.force_disconnect_client(client_id) break await self._redis.srem(USER_CONN_KEY.format(user_id=user_id), client_id) else: # Multi-device — reject so caller can show blocked image raise MaxConnectionsExceededError( f"User {user_id} already has {count}/{max_conns} connections" ) def _sort_by_preference( self, provider_maps: list[dict], preferred_ids: list[int] | None ) -> list[dict]: if not preferred_ids: return provider_maps pref_set = set(preferred_ids) preferred = [pm for pm in provider_maps if pm["provider_account_id"] in pref_set] rest = [pm for pm in provider_maps if pm["provider_account_id"] not in pref_set] return preferred + rest async def _build_stream_urls(self, account_id: int, primary_stream_url: str) -> list[str]: """Build list of stream URLs for multi-domain failover. Only includes domains with status != 'error' (healthy or not-yet-checked). If ALL domains are currently marked 'error', falls back to using all of them so the stream isn't stranded — the health checker will restore them when reachable again. """ from urllib.parse import urlparse, urlunparse from ..models.provider import ProviderUrl from ..database import AsyncSessionLocal from sqlalchemy import select parsed = urlparse(primary_stream_url) async with AsyncSessionLocal() as db: # Prefer healthy/unknown domains only result = await db.execute( select(ProviderUrl) .where( ProviderUrl.provider_account_id == account_id, ProviderUrl.is_active == True, # noqa: E712 ProviderUrl.status != "error", ) .order_by(ProviderUrl.priority) ) provider_urls = result.scalars().all() if not provider_urls: # Fallback: all active domains even if all are 'error' result = await db.execute( select(ProviderUrl) .where( ProviderUrl.provider_account_id == account_id, ProviderUrl.is_active == True, # noqa: E712 ) .order_by(ProviderUrl.priority) ) provider_urls = result.scalars().all() if provider_urls: logger.warning( f"provider {account_id}: all domains marked 'error' — " f"using all {len(provider_urls)} as fallback" ) if not provider_urls: return [primary_stream_url] skipped = [] urls = [] for pu in provider_urls: alt = urlparse(pu.url.rstrip('/')) reconstructed = urlunparse(( alt.scheme or parsed.scheme, alt.netloc, parsed.path, parsed.params, parsed.query, parsed.fragment, )) urls.append(reconstructed) if skipped: logger.debug(f"provider {account_id}: skipped error domains {skipped}") return urls async def _find_free_slot(self, provider_maps: list[dict]) -> tuple[int | None, str | None]: for pm in provider_maps: account_id = pm["provider_account_id"] slot_key = PROVIDER_SLOT_KEY.format(account_id=account_id) busy = await self._redis.exists(slot_key) if not busy: return account_id, pm["stream_url"] return None, None async def _find_eviction_candidate( self, provider_maps: list[dict] ) -> tuple[int | None, str | None, BroadcastGroup | None]: """Find a stream with only non-priority clients that uses a slot we need.""" available = {pm["provider_account_id"]: pm["stream_url"] for pm in provider_maps} candidate: BroadcastGroup | None = None for group in self._groups.values(): if group.provider_account_id in available and group.all_non_priority: if candidate is None or group.client_count < candidate.client_count: candidate = group if candidate is None: return None, None, None return candidate.provider_account_id, available[candidate.provider_account_id], candidate async def _publish_monitoring(self) -> None: try: data = json.dumps([g.stats() for g in self._groups.values()]) await self._redis.publish(MONITORING_KEY, data) except Exception: pass pool = ProviderPool()