from __future__ import annotations from dataclasses import dataclass from datetime import datetime from typing import Any, NamedTuple, Self, TypedDict, cast import sqlalchemy as sa from fansifter_common.adapters.db.models import Query from fansifter_common.adapters.db.types import SafeJSONType from pydantic import TypeAdapter from sqlalchemy.dialects.postgresql import JSONB, insert from sqlalchemy.orm import Mapped, mapped_column from app.adapters.db import Model from app.dsp.enums import DSPId from app.dsp.models import DSPClient from app.fandata.enums import ( CollectionGranularity, FanCollectionStatus, FanCredentialsStatus, ) from app.fandata.types import FanStepState # ─── Fan credentials ───────────────────────────────────────────────────────── class FanCredentialsRow(TypedDict): dsp_user_id: str dsp_id: DSPId client_id: str email: str refresh_token_encrypted: str class FanViewRow(TypedDict): dsp_user_id: str dsp_id: DSPId token_status: FanCredentialsStatus first_seen_at: datetime last_collected_at: datetime | None last_collection_status: FanCollectionStatus | None consecutive_failures: int dsp_client_id: int | None dsp_client_name: str | None class FanViewPaginated(NamedTuple): items: list[FanViewRow] total: int class FanCredentialsQuery(Query["FanCredentials"]): def filter(self, f: FanCredentialsFilter) -> Self: self._query = f.apply(sa.select(FanCredentials)) return self def count(self) -> int: count_stmt = self._query.with_only_columns( sa.func.count(), maintain_column_froms=True ) return self.session.scalar(count_stmt) or 0 def cursor_page(self, *, after_id: int, limit: int) -> list[FanCredentials]: stmt = ( self._query.where(FanCredentials.id > after_id) .order_by(FanCredentials.id) .limit(limit) ) return list(self.session.execute(stmt).scalars()) def upsert(self, rows: list[FanCredentialsRow]) -> None: stmt = insert(FanCredentials).on_conflict_do_update( index_elements=["dsp_user_id", "dsp_id", "client_id"], set_={ "email": insert(FanCredentials).excluded.email, "refresh_token_encrypted": insert( FanCredentials ).excluded.refresh_token_encrypted, "token_status": FanCredentialsStatus.active, "token_updated_at": sa.func.now(), }, ) self.session.execute(stmt, rows) def mark_revoked(self, dsp_user_id: str, dsp_id: DSPId, client_id: str) -> None: self.session.execute( sa.update(FanCredentials) .where( FanCredentials.dsp_user_id == dsp_user_id, FanCredentials.dsp_id == dsp_id, FanCredentials.client_id == client_id, ) .values(token_status=FanCredentialsStatus.revoked) ) def update_refresh_token( self, dsp_user_id: str, dsp_id: DSPId, client_id: str, refresh_token_encrypted: str, ) -> None: self.session.execute( sa.update(FanCredentials) .where( FanCredentials.dsp_user_id == dsp_user_id, FanCredentials.dsp_id == dsp_id, FanCredentials.client_id == client_id, ) .values( refresh_token_encrypted=refresh_token_encrypted, token_updated_at=sa.func.now(), ) ) def fan_view_paginate( self, *, dsp_id: DSPId | None = None, token_status: FanCredentialsStatus | None = None, search: str | None = None, limit: int = 20, offset: int = 0, ) -> FanViewPaginated: # One row per credential (dsp_user_id, dsp_id, client_id). stmt = ( sa.select( FanCredentials.dsp_user_id, FanCredentials.dsp_id, FanCredentials.token_status, FanCredentials.created_at.label("first_seen_at"), FanCollectionState.last_collected_at, FanCollectionState.last_status.label("last_collection_status"), sa.func.coalesce(FanCollectionState.consecutive_failures, 0).label( "consecutive_failures" ), DSPClient.id.label("dsp_client_id"), DSPClient.display_name.label("dsp_client_name"), ) .outerjoin( FanCollectionState, sa.and_( FanCredentials.dsp_user_id == FanCollectionState.dsp_user_id, FanCredentials.dsp_id == FanCollectionState.dsp_id, ), ) .outerjoin( DSPClient, sa.and_( FanCredentials.dsp_id == DSPClient.dsp_id, FanCredentials.client_id == DSPClient.client_id, ), ) ) if dsp_id is not None: stmt = stmt.where(FanCredentials.dsp_id == dsp_id) if token_status is not None: stmt = stmt.where(FanCredentials.token_status == token_status) if search: stmt = stmt.where(FanCredentials.dsp_user_id.ilike(f"%{search}%")) count_sq = stmt.subquery() total = ( self.session.scalar(sa.select(sa.func.count()).select_from(count_sq)) or 0 ) rows = ( self.session.execute( stmt.order_by(FanCredentials.dsp_user_id, FanCredentials.dsp_id) .limit(limit) .offset(offset) ) .mappings() .all() ) return FanViewPaginated(items=cast(list[FanViewRow], list(rows)), total=total) class FanCredentials(Model, kw_only=True): __tablename__ = "fan_credentials" __table_args__ = (sa.UniqueConstraint("dsp_user_id", "dsp_id", "client_id"),) id: Mapped[int] = mapped_column(primary_key=True) dsp_user_id: Mapped[str] dsp_id: Mapped[DSPId] client_id: Mapped[str] email: Mapped[str] refresh_token_encrypted: Mapped[str] token_status: Mapped[FanCredentialsStatus] = mapped_column( default=FanCredentialsStatus.active ) token_updated_at: Mapped[datetime] = mapped_column(server_default=sa.func.now()) created_at: Mapped[datetime] = mapped_column(server_default=sa.func.now()) query = FanCredentialsQuery.as_descriptor() @dataclass(kw_only=True, frozen=True) class FanCredentialsFilter: dsp_client_id: int token_status: FanCredentialsStatus | None = FanCredentialsStatus.active not_collected_since: datetime | None = None max_consecutive_failures: int | None = None limit: int | None = None def apply[S: tuple[Any, ...]](self, stmt: sa.Select[S]) -> sa.Select[S]: if self.dsp_client_id is not None: stmt = stmt.join( DSPClient, sa.and_( FanCredentials.client_id == DSPClient.client_id, FanCredentials.dsp_id == DSPClient.dsp_id, DSPClient.id == self.dsp_client_id, ), ) if ( self.not_collected_since is not None or self.max_consecutive_failures is not None ): stmt = stmt.outerjoin( FanCollectionState, sa.and_( FanCredentials.dsp_user_id == FanCollectionState.dsp_user_id, FanCredentials.dsp_id == FanCollectionState.dsp_id, ), ) if self.token_status is not None: stmt = stmt.where(FanCredentials.token_status == self.token_status) if self.not_collected_since is not None: stmt = stmt.where( sa.or_( FanCollectionState.last_collected_at.is_(None), FanCollectionState.last_collected_at < self.not_collected_since, ) ) if self.max_consecutive_failures is not None: stmt = stmt.where( sa.or_( FanCollectionState.consecutive_failures.is_(None), FanCollectionState.consecutive_failures <= self.max_consecutive_failures, ) ) return stmt def dump_python(self) -> dict[str, Any]: return _fan_credentials_filter_adapter.dump_python( self, mode="json", exclude_none=True, exclude={"dsp_client_id"} ) @classmethod def validate_python(cls, self) -> FanCredentialsFilter: return _fan_credentials_filter_adapter.validate_python(self) _fan_credentials_filter_adapter = TypeAdapter(FanCredentialsFilter) # ─── Fan collection state ───────────────────────────────────────────────────── class FanCollectionStateRow(NamedTuple): date: str success: int errors: int throttled: int = 0 # Lightweight table reference to avoid circular imports with pipeline models _pipeline_run = sa.table( "pipeline_run", sa.column("started_at"), sa.column("rate_limited", sa.Integer), ) _EMPTY_STEP_STATE = FanStepState() class FanCollectionStateQuery(Query["FanCollectionState"]): def get_step_state(self, *, dsp_user_id: str, dsp_id: DSPId) -> FanStepState: row = self.session.execute( sa.select(FanCollectionState.step_state).where( FanCollectionState.dsp_user_id == dsp_user_id, FanCollectionState.dsp_id == dsp_id, ) ).scalar_one_or_none() if row is None: return _EMPTY_STEP_STATE return FanStepState.model_validate(row) def record_step( self, *, dsp_user_id: str, dsp_id: DSPId, dsp_client_id: int, step: str, collected_at: datetime, ) -> None: patch = {f"{step}_collected_at": collected_at.isoformat()} stmt = ( insert(FanCollectionState) .values( dsp_user_id=dsp_user_id, dsp_id=dsp_id, dsp_client_id=dsp_client_id, consecutive_failures=0, step_state=patch, ) .on_conflict_do_update( index_elements=["dsp_user_id", "dsp_id"], set_={ "step_state": sa.cast( sa.cast(FanCollectionState.step_state, JSONB).op("||")(patch), sa.Text, ) }, ) ) self.session.execute(stmt) def recently_collected_ids( self, *, dsp_user_ids: list[str], dsp_id: DSPId, since: datetime, ) -> frozenset[str]: """Return the subset of dsp_user_ids collected at or after `since`.""" if not dsp_user_ids: return frozenset() rows = self.session.execute( sa.select(FanCollectionState.dsp_user_id).where( FanCollectionState.dsp_user_id.in_(dsp_user_ids), FanCollectionState.dsp_id == dsp_id, FanCollectionState.last_collected_at >= since, ) ).scalars() return frozenset(rows) def record( self, *, dsp_user_id: str, dsp_id: DSPId, dsp_client_id: int, status: FanCollectionStatus, collected_at: datetime, ) -> None: is_failure = status not in ( FanCollectionStatus.success, FanCollectionStatus.rate_limited, ) stmt = insert(FanCollectionState).values( dsp_user_id=dsp_user_id, dsp_id=dsp_id, dsp_client_id=dsp_client_id, last_collected_at=collected_at, last_status=status, consecutive_failures=1 if is_failure else 0, ) stmt = stmt.on_conflict_do_update( index_elements=["dsp_user_id", "dsp_id"], set_={ "dsp_client_id": stmt.excluded.dsp_client_id, "last_collected_at": stmt.excluded.last_collected_at, "last_status": stmt.excluded.last_status, "consecutive_failures": ( sa.func.coalesce(FanCollectionState.consecutive_failures, 0) + 1 if is_failure else 0 ), }, ) self.session.execute(stmt) def collection_activity( self, *, days: int = 7, granularity: CollectionGranularity = CollectionGranularity.daily, ) -> list[FanCollectionStateRow]: cutoff = sa.func.now() - sa.text(f"interval '{days} days'") if granularity == CollectionGranularity.hourly: fan_bucket_expr = sa.func.date_trunc( "hour", FanCollectionState.last_collected_at ) run_bucket_expr = sa.func.date_trunc("hour", _pipeline_run.c.started_at) else: fan_bucket_expr = sa.func.date(FanCollectionState.last_collected_at) run_bucket_expr = sa.func.date(_pipeline_run.c.started_at) # Per-bucket fan success/error counts from fan_collection_state fan_q = ( sa.select( fan_bucket_expr.label("bucket"), sa.func.count() .filter(FanCollectionState.last_status == FanCollectionStatus.success) .label("success"), sa.func.count() .filter(FanCollectionState.last_status != FanCollectionStatus.success) .label("errors"), ) .where(FanCollectionState.last_collected_at >= cutoff) .group_by(fan_bucket_expr) .subquery("fan_activity") ) # Per-bucket throttled counts from pipeline_run run_q = ( sa.select( run_bucket_expr.label("bucket"), sa.func.sum(_pipeline_run.c.rate_limited).label("throttled"), ) .where(_pipeline_run.c.started_at >= cutoff) .where(_pipeline_run.c.rate_limited > 0) .group_by(run_bucket_expr) .subquery("run_activity") ) stmt = ( sa.select( fan_q.c.bucket, fan_q.c.success, fan_q.c.errors, sa.func.coalesce(run_q.c.throttled, 0).label("throttled"), ) .select_from(fan_q.outerjoin(run_q, fan_q.c.bucket == run_q.c.bucket)) .order_by(fan_q.c.bucket) ) rows = self.session.execute(stmt).all() return [ FanCollectionStateRow( date=str(row.bucket), success=row.success, errors=row.errors, throttled=row.throttled, ) for row in rows ] class FanCollectionState(Model, kw_only=True): __tablename__ = "fan_collection_state" dsp_user_id: Mapped[str] = mapped_column(primary_key=True) dsp_id: Mapped[DSPId] = mapped_column(primary_key=True) dsp_client_id: Mapped[int | None] = mapped_column(default=None) last_collected_at: Mapped[datetime | None] = mapped_column(default=None) last_status: Mapped[FanCollectionStatus | None] = mapped_column(default=None) consecutive_failures: Mapped[int] = mapped_column(default=0) step_state: Mapped[dict[str, Any]] = mapped_column( SafeJSONType, server_default="'{}'" ) query = FanCollectionStateQuery.as_descriptor() # ─── Fan ───────────────────────────────────────────────────────────────────── class FanRow(TypedDict): dsp_user_id: str dsp_id: DSPId email: str fan_id: str country: str | None product: str | None display_name: str | None collected_at: datetime class FanQuery(Query["Fan"]): def upsert(self, rows: list[FanRow]) -> None: stmt = insert(Fan) stmt = stmt.on_conflict_do_update( index_elements=["dsp_user_id", "dsp_id"], set_={ "email": stmt.excluded.email, "fan_id": stmt.excluded.fan_id, "country": stmt.excluded.country, "product": stmt.excluded.product, "display_name": stmt.excluded.display_name, "collected_at": stmt.excluded.collected_at, }, ) self.session.execute(stmt, rows) def collected_today(self) -> int: today = sa.func.date_trunc("day", sa.func.now()) return self.where(Fan.collected_at >= today).count() class Fan(Model, kw_only=True): __tablename__ = "sink_fan" dsp_user_id: Mapped[str] = mapped_column(primary_key=True) dsp_id: Mapped[DSPId] = mapped_column(primary_key=True) email: Mapped[str] fan_id: Mapped[str] country: Mapped[str | None] = mapped_column(default=None) product: Mapped[str | None] = mapped_column(default=None) display_name: Mapped[str | None] = mapped_column(default=None) collected_at: Mapped[datetime] query = FanQuery.as_descriptor() # ─── Fan top artist ─────────────────────────────────────────────────────────── class FanTopArtistRow(TypedDict): dsp_id: DSPId dsp_user_id: str artist_id: str data: dict[str, Any] collected_at: datetime class FanTopArtistQuery(Query["FanTopArtist"]): def upsert(self, rows: list[FanTopArtistRow]) -> None: stmt = insert(FanTopArtist) stmt = stmt.on_conflict_do_update( index_elements=["dsp_id", "dsp_user_id", "artist_id"], set_={ "data": stmt.excluded.data, "collected_at": stmt.excluded.collected_at, }, ) self.session.execute(stmt, rows) class FanTopArtist(Model, kw_only=True): __tablename__ = "sink_fan_top_artist" dsp_id: Mapped[DSPId] = mapped_column(primary_key=True) dsp_user_id: Mapped[str] = mapped_column(primary_key=True) artist_id: Mapped[str] = mapped_column(primary_key=True) data: Mapped[dict[str, Any]] = mapped_column(JSONB) collected_at: Mapped[datetime] query = FanTopArtistQuery.as_descriptor() # ─── Fan recently played ────────────────────────────────────────────────────── class FanRecentlyPlayedRow(TypedDict): dsp_id: DSPId dsp_user_id: str track_id: str played_at: datetime data: dict[str, Any] collected_at: datetime class FanRecentlyPlayedQuery(Query["FanRecentlyPlayed"]): def upsert(self, rows: list[FanRecentlyPlayedRow]) -> None: stmt = insert(FanRecentlyPlayed) stmt = stmt.on_conflict_do_update( index_elements=["dsp_id", "dsp_user_id", "track_id"], set_={ "played_at": stmt.excluded.played_at, "data": stmt.excluded.data, "collected_at": stmt.excluded.collected_at, }, ) self.session.execute(stmt, rows) class FanRecentlyPlayed(Model, kw_only=True): __tablename__ = "sink_fan_recently_played" dsp_id: Mapped[DSPId] = mapped_column(primary_key=True) dsp_user_id: Mapped[str] = mapped_column(primary_key=True) track_id: Mapped[str] = mapped_column(primary_key=True) played_at: Mapped[datetime] data: Mapped[dict[str, Any]] = mapped_column(JSONB) collected_at: Mapped[datetime] query = FanRecentlyPlayedQuery.as_descriptor() # ─── Fan top track ──────────────────────────────────────────────────────────── class FanTopTrackRow(TypedDict): dsp_id: DSPId dsp_user_id: str track_id: str data: dict[str, Any] collected_at: datetime class FanTopTrackQuery(Query["FanTopTrack"]): def upsert(self, rows: list[FanTopTrackRow]) -> None: stmt = insert(FanTopTrack) stmt = stmt.on_conflict_do_update( index_elements=["dsp_id", "dsp_user_id", "track_id"], set_={ "data": stmt.excluded.data, "collected_at": stmt.excluded.collected_at, }, ) self.session.execute(stmt, rows) class FanTopTrack(Model, kw_only=True): __tablename__ = "sink_fan_top_track" dsp_id: Mapped[DSPId] = mapped_column(primary_key=True) dsp_user_id: Mapped[str] = mapped_column(primary_key=True) track_id: Mapped[str] = mapped_column(primary_key=True) data: Mapped[dict[str, Any]] = mapped_column(JSONB) collected_at: Mapped[datetime] query = FanTopTrackQuery.as_descriptor() # ─── Fan playlist ───────────────────────────────────────────────────────────── class FanPlaylistRow(TypedDict): dsp_id: DSPId dsp_user_id: str playlist_id: str data: dict[str, Any] collected_at: datetime class FanPlaylistQuery(Query["FanPlaylist"]): def upsert(self, rows: list[FanPlaylistRow]) -> None: stmt = insert(FanPlaylist) stmt = stmt.on_conflict_do_update( index_elements=["dsp_id", "dsp_user_id", "playlist_id"], set_={ "data": stmt.excluded.data, "collected_at": stmt.excluded.collected_at, }, ) self.session.execute(stmt, rows) class FanPlaylist(Model, kw_only=True): __tablename__ = "sink_fan_playlist" dsp_id: Mapped[DSPId] = mapped_column(primary_key=True) dsp_user_id: Mapped[str] = mapped_column(primary_key=True) playlist_id: Mapped[str] = mapped_column(primary_key=True) data: Mapped[dict[str, Any]] = mapped_column(JSONB) collected_at: Mapped[datetime] query = FanPlaylistQuery.as_descriptor() # ─── Fan saved album ────────────────────────────────────────────────────────── class FanSavedAlbumRow(TypedDict): dsp_id: DSPId dsp_user_id: str album_id: str added_at: datetime data: dict[str, Any] collected_at: datetime class FanSavedAlbumQuery(Query["FanSavedAlbum"]): def upsert(self, rows: list[FanSavedAlbumRow]) -> None: stmt = insert(FanSavedAlbum) stmt = stmt.on_conflict_do_update( index_elements=["dsp_id", "dsp_user_id", "album_id"], set_={ "added_at": stmt.excluded.added_at, "data": stmt.excluded.data, "collected_at": stmt.excluded.collected_at, }, ) self.session.execute(stmt, rows) class FanSavedAlbum(Model, kw_only=True): __tablename__ = "sink_fan_saved_album" dsp_id: Mapped[DSPId] = mapped_column(primary_key=True) dsp_user_id: Mapped[str] = mapped_column(primary_key=True) album_id: Mapped[str] = mapped_column(primary_key=True) added_at: Mapped[datetime] data: Mapped[dict[str, Any]] = mapped_column(JSONB) collected_at: Mapped[datetime] query = FanSavedAlbumQuery.as_descriptor() # ─── Fan saved track ────────────────────────────────────────────────────────── class FanSavedTrackRow(TypedDict): dsp_id: DSPId dsp_user_id: str track_id: str added_at: datetime data: dict[str, Any] collected_at: datetime class FanSavedTrackQuery(Query["FanSavedTrack"]): def upsert(self, rows: list[FanSavedTrackRow]) -> None: stmt = insert(FanSavedTrack) stmt = stmt.on_conflict_do_update( index_elements=["dsp_id", "dsp_user_id", "track_id"], set_={ "added_at": stmt.excluded.added_at, "data": stmt.excluded.data, "collected_at": stmt.excluded.collected_at, }, ) self.session.execute(stmt, rows) class FanSavedTrack(Model, kw_only=True): __tablename__ = "sink_fan_saved_track" dsp_id: Mapped[DSPId] = mapped_column(primary_key=True) dsp_user_id: Mapped[str] = mapped_column(primary_key=True) track_id: Mapped[str] = mapped_column(primary_key=True) added_at: Mapped[datetime] data: Mapped[dict[str, Any]] = mapped_column(JSONB) collected_at: Mapped[datetime] query = FanSavedTrackQuery.as_descriptor() # ─── Fan followed artist ────────────────────────────────────────────────────── class FanFollowedArtistRow(TypedDict): dsp_id: DSPId dsp_user_id: str artist_id: str data: dict[str, Any] collected_at: datetime class FanFollowedArtistQuery(Query["FanFollowedArtist"]): def upsert(self, rows: list[FanFollowedArtistRow]) -> None: stmt = insert(FanFollowedArtist) stmt = stmt.on_conflict_do_update( index_elements=["dsp_id", "dsp_user_id", "artist_id"], set_={ "data": stmt.excluded.data, "collected_at": stmt.excluded.collected_at, }, ) self.session.execute(stmt, rows) class FanFollowedArtist(Model, kw_only=True): __tablename__ = "sink_fan_followed_artist" dsp_id: Mapped[DSPId] = mapped_column(primary_key=True) dsp_user_id: Mapped[str] = mapped_column(primary_key=True) artist_id: Mapped[str] = mapped_column(primary_key=True) data: Mapped[dict[str, Any]] = mapped_column(JSONB) collected_at: Mapped[datetime] query = FanFollowedArtistQuery.as_descriptor()