import logging import random import time from collections.abc import Callable from typing import Any, Literal from fansifter_common.httpclient import HTTPClientError from fansifter_common.utils.functional import lazy_proxy from pydantic import SecretStr, ValidationError from resonance_engine.adapters import aws_secretsmanager from resonance_engine.adapters.amazon_music import AmazonMusicClient from resonance_engine.adapters.apple_music import AppleMusicClient from resonance_engine.adapters.deezer import DeezerClient from resonance_engine.adapters.redis import redis_client from resonance_engine.adapters.spotify import SpotifyClient from resonance_engine.config import settings from resonance_engine.dsp.enums import DSPClientName, DSPId, DSPResource from .backends import ( AmazonMusicDSPBackend, AppleMusicDSPBackend, DeezerDSPBackend, DSPBackend, SpotifyDSPBackend, ) from .exceptions import ( DSPForbiddenError, DSPNotRegisteredError, DSPResourceUnsupportedError, InvalidResponseError, RateLimitError, StreamingAPIError, ) from .stats import InMemoryStatsBackend, RedisStatsBackend, RequestStats, StatsBackend from .types import TokenResult logger = logging.getLogger(__name__) DEFAULT_RETRY_AFTER = 5 DEFAULT_MAX_RETRIES = 3 RETRY_AFTER_THRESHOLD = 10 # seconds — escalate to RateLimitError immediately MAX_SLEEP = 30 TRANSIENT_STATUSES = frozenset({502, 503, 504}) TRANSIENT_BACKOFF_S = 0.5 class DSPGateway: def __init__( self, max_retries: int = DEFAULT_MAX_RETRIES, default_retry_after: int = DEFAULT_RETRY_AFTER, stats_backend: Literal["memory", "redis"] = "memory", ) -> None: self._backends: dict[DSPClientName, DSPBackend] = {} self._stats: dict[DSPClientName, StatsBackend] = {} self._max_retries = max_retries self._default_retry_after = default_retry_after self._stats_backend_kind: Literal["memory", "redis"] = stats_backend @staticmethod def _create_stats_backend( kind: Literal["memory", "redis"], client_name: DSPClientName ) -> StatsBackend: if kind == "redis": return RedisStatsBackend(redis_client, f"resonance:streaming:{client_name}") return InMemoryStatsBackend() def register( self, client_name: DSPClientName, backend: DSPBackend, ) -> None: self._backends[client_name] = backend def is_configured(self, client_name: DSPClientName) -> bool: return client_name in self._backends def get(self, client_name: DSPClientName) -> DSPBackend: try: return self._backends[client_name] except KeyError: raise DSPNotRegisteredError(client_name) from None def stats(self, client_name: DSPClientName) -> RequestStats: return self._stats_for(client_name).snapshot() def _stats_for(self, client_name: DSPClientName) -> StatsBackend: if client_name not in self._stats: self._stats[client_name] = self._create_stats_backend( self._stats_backend_kind, client_name ) return self._stats[client_name] def _call[T]( self, client_name: DSPClientName, fn: Callable[[], T], *, resource: DSPResource | None = None, scope: str | None = None, ) -> T: backend = self.get(client_name) if resource is not None and not backend.has_required_scope(resource, scope): # Skip the request entirely — the token can't access this endpoint. raise DSPForbiddenError( f"token missing scope '{backend.required_scope(resource)}' for {resource}", status_code=403, request_made=False, ) def wrapped() -> T: try: return fn() except HTTPClientError as exc: raise StreamingAPIError.from_http_error(exc) from exc except ValidationError as exc: raise InvalidResponseError.from_validation_error(exc) from exc return self._execute_with_retry(client_name, wrapped) def _execute_with_retry[T]( self, client_name: DSPClientName, fn: Callable[[], T] ) -> T: stats = self._stats_for(client_name) counted = False def count_once() -> None: # Count one request per logical call (not per retry), and only when one # actually reached the DSP — an unsupported resource makes no request. nonlocal counted if not counted: stats.incr_requests() counted = True for attempt in range(self._max_retries + 1): try: result = fn() except DSPResourceUnsupportedError: raise except StreamingAPIError as exc: count_once() if exc.status_code == 429: self._handle_rate_limit(stats, exc, attempt, client_name) continue if exc.status_code in TRANSIENT_STATUSES: if attempt == self._max_retries: raise stats.incr_retries() sleep = TRANSIENT_BACKOFF_S * (2**attempt) logger.warning( "Transient %d, retrying in %.1fs (attempt %d/%d)", exc.status_code, sleep, attempt + 1, self._max_retries, exc_info=exc, extra={ "dsp_client_name": client_name, }, ) time.sleep(sleep) continue raise except Exception: count_once() raise count_once() return result raise AssertionError("unreachable") def _handle_rate_limit( self, stats: StatsBackend, exc: StreamingAPIError, attempt: int, client_name: DSPClientName, ) -> None: stats.incr_rate_limited() retry_after = exc.retry_after or self._default_retry_after if retry_after > RETRY_AFTER_THRESHOLD: raise RateLimitError( f"Retry-After {retry_after}s exceeds threshold " f"({RETRY_AFTER_THRESHOLD}s) — aborting batch", status_code=429, retry_after=retry_after, ) from exc if attempt == self._max_retries: raise RateLimitError( f"Rate limit retries exhausted after {attempt + 1} attempts", status_code=429, retry_after=retry_after, ) from exc stats.incr_retries() sleep = min(retry_after * (2**attempt), MAX_SLEEP) + random.uniform(0, 1) logger.warning( "Rate limited (429), retrying in %.1fs (attempt %d/%d)", sleep, attempt + 1, self._max_retries, extra={ "dsp_client_name": client_name, }, ) time.sleep(sleep) # -- API ----------------------------------------------------------------- def refresh_token( self, client_name: DSPClientName, refresh_token: SecretStr ) -> TokenResult: return self._call( client_name, lambda: self.get(client_name).refresh_token(refresh_token) ) def get_profile( self, client_name: DSPClientName, access_token: SecretStr, *, scope: str | None = None, ) -> dict[str, Any]: return self._call( client_name, lambda: self.get(client_name).get_profile(access_token), resource=DSPResource.profile, scope=scope, ) def get_top_artists( self, client_name: DSPClientName, access_token: SecretStr, *, scope: str | None = None, ) -> list[dict[str, Any]]: return self._call( client_name, lambda: self.get(client_name).get_top_artists(access_token), resource=DSPResource.top_artists, scope=scope, ) def get_top_tracks( self, client_name: DSPClientName, access_token: SecretStr, *, scope: str | None = None, ) -> list[dict[str, Any]]: return self._call( client_name, lambda: self.get(client_name).get_top_tracks(access_token), resource=DSPResource.top_tracks, scope=scope, ) def get_recently_played( self, client_name: DSPClientName, access_token: SecretStr, *, after: int | None = None, scope: str | None = None, ) -> list[dict[str, Any]]: return self._call( client_name, lambda: self.get(client_name).get_recently_played( access_token, after=after ), resource=DSPResource.recently_played, scope=scope, ) def get_playlists( self, client_name: DSPClientName, access_token: SecretStr, *, scope: str | None = None, ) -> list[dict[str, Any]]: return self._call( client_name, lambda: self.get(client_name).get_playlists(access_token), resource=DSPResource.playlists, scope=scope, ) def get_saved_albums( self, client_name: DSPClientName, access_token: SecretStr, *, scope: str | None = None, ) -> list[dict[str, Any]]: return self._call( client_name, lambda: self.get(client_name).get_saved_albums(access_token), resource=DSPResource.saved_albums, scope=scope, ) def get_saved_tracks( self, client_name: DSPClientName, access_token: SecretStr, *, scope: str | None = None, ) -> list[dict[str, Any]]: return self._call( client_name, lambda: self.get(client_name).get_saved_tracks(access_token), resource=DSPResource.saved_tracks, scope=scope, ) def get_followed_artists( self, client_name: DSPClientName, access_token: SecretStr, *, scope: str | None = None, ) -> list[dict[str, Any]]: return self._call( client_name, lambda: self.get(client_name).get_followed_artists(access_token), resource=DSPResource.followed_artists, scope=scope, ) def close(self) -> None: for backend in self._backends.values(): backend.close() def __enter__(self) -> DSPGateway: return self def __exit__(self, *_: object) -> None: self.close() # --------------------------------------------------------------------------- # Registration # --------------------------------------------------------------------------- def get_dsp_gateway() -> DSPGateway: gateway = DSPGateway(stats_backend=settings.dsp_stats_backend) for name, creds in settings.dsp_clients.items(): backend: DSPBackend match creds["dsp_id"]: case DSPId.spotify: backend = SpotifyDSPBackend( client=SpotifyClient( client_id=creds["client_id"], client_secret=aws_secretsmanager.get_secret_or_default( creds["client_secret_name"], default=creds["client_secret"], ), ) ) case DSPId.deezer: backend = DeezerDSPBackend(client=DeezerClient()) case DSPId.amazon: backend = AmazonMusicDSPBackend( client=AmazonMusicClient( client_id=creds["client_id"], client_secret=aws_secretsmanager.get_secret_or_default( creds["client_secret_name"], default=creds["client_secret"], ), profile_id=creds["profile_id"], ) ) case DSPId.apple: backend = AppleMusicDSPBackend( client=AppleMusicClient( team_id=creds["team_id"], key_id=creds["key_id"], private_key=aws_secretsmanager.get_secret_or_default( creds["private_key_name"], default=creds["private_key"], ), ) ) case _: continue gateway.register(name, backend=backend) return gateway dsp_gateway = lazy_proxy(get_dsp_gateway)