from __future__ import annotations import uuid from collections.abc import Sequence from datetime import datetime, timedelta from typing import NamedTuple, Self import sqlalchemy as sa from fansifter_common.adapters.db.models import Query from fansifter_common.adapters.db.types import PydanticType from fansifter_common.utils import timezone from sqlalchemy.orm import Mapped, mapped_column, relationship from app.adapters.db import Model from app.config import settings from app.dsp.models import DSPClient from app.fandata.models import FanCredentialsFilter from app.pipeline.enums import RunFinishedReason, RunSource, RunStatus class PipelineRunPaginated(NamedTuple): items: Sequence[PipelineRun] total: int class PipelineRunQuery(Query["PipelineRun"]): def active(self) -> Self: return self.where( PipelineRun.status.in_( [ RunStatus.queued, RunStatus.running, RunStatus.paused, ], ) ) def paginate(self, limit: int, offset: int) -> PipelineRunPaginated: items = ( self.order_by(PipelineRun.started_at.desc()) .limit(limit) .offset(offset) .all() ) total = ( self.session.scalar(sa.select(sa.func.count()).select_from(self.model)) or 0 ) return PipelineRunPaginated(items=items, total=total) def increment_batch_result( self, run_id: uuid.UUID | str, *, processed: int, errors: int, rate_limited: int, stale_tokens: int, steps_skipped: int = 0, api_requests: int = 0, ) -> tuple[int, int]: """Increment batch metrics and return (completed_batches, total_batches).""" row = self.session.execute( sa.update(PipelineRun) .where(PipelineRun.id == run_id) .values( fans_processed=sa.func.coalesce(PipelineRun.fans_processed, 0) + processed, errors=sa.func.coalesce(PipelineRun.errors, 0) + errors, rate_limited=sa.func.coalesce(PipelineRun.rate_limited, 0) + rate_limited, stale_tokens=sa.func.coalesce(PipelineRun.stale_tokens, 0) + stale_tokens, steps_skipped=sa.func.coalesce(PipelineRun.steps_skipped, 0) + steps_skipped, total_requests=sa.func.coalesce(PipelineRun.total_requests, 0) + api_requests, completed_batches=PipelineRun.completed_batches + 1, ) .returning(PipelineRun.completed_batches, PipelineRun.total_batches) ).one_or_none() if row is None: return 0, 0 return row.completed_batches, row.total_batches class PipelineRun(Model, kw_only=True): __tablename__ = "pipeline_run" __table_args__ = ( sa.Index( "pipeline_run_one_active_per_dsp_client", "dsp_client_id", unique=True, postgresql_where=sa.text("status IN ('queued', 'running', 'paused')"), ), ) id: Mapped[uuid.UUID] = mapped_column( primary_key=True, default_factory=uuid.uuid4, server_default=sa.func.gen_random_uuid(), ) status: Mapped[RunStatus] = mapped_column(default=RunStatus.idle) dsp_client_id: Mapped[int] = mapped_column(sa.Integer) source: Mapped[RunSource | None] = mapped_column(default=None) total_fans: Mapped[int] = mapped_column(default=0) total_batches: Mapped[int] = mapped_column(default=0) fans_processed: Mapped[int] = mapped_column(default=0) completed_batches: Mapped[int] = mapped_column(default=0) errors: Mapped[int] = mapped_column(default=0) rate_limited: Mapped[int] = mapped_column(default=0) stale_tokens: Mapped[int] = mapped_column(default=0) steps_skipped: Mapped[int] = mapped_column(default=0) total_requests: Mapped[int] = mapped_column(default=0) filters: Mapped[FanCredentialsFilter | None] = mapped_column( PydanticType(FanCredentialsFilter | None), default=None ) created_at: Mapped[datetime] = mapped_column( default_factory=timezone.now, server_default=sa.func.now() ) started_at: Mapped[datetime | None] = mapped_column(default=None) finished_at: Mapped[datetime | None] = mapped_column(default=None) finished_reason: Mapped[RunFinishedReason | None] = mapped_column(default=None) # Relations dsp_client: Mapped[DSPClient] = relationship( lazy="joined", init=False, primaryjoin=lambda: PipelineRun.dsp_client_id == DSPClient.id, foreign_keys=lambda: [PipelineRun.dsp_client_id], ) query = PipelineRunQuery.as_descriptor() @property def is_stale(self) -> bool: if self.status not in (RunStatus.queued, RunStatus.running, RunStatus.paused): return False if not self.started_at: return False return self.started_at < timezone.now() - timedelta( seconds=settings.pipeline_run_stale_timeout_s ) @property def eta_seconds(self) -> int | None: if self.status not in (RunStatus.running, RunStatus.paused): return None if not self.started_at or self.fans_processed <= 0 or self.total_fans <= 0: return None if self.fans_processed >= self.total_fans: return None elapsed_s = (timezone.now() - self.started_at).total_seconds() if elapsed_s <= 0: return None rate = self.fans_processed / elapsed_s return round((self.total_fans - self.fans_processed) / rate) @property def throughput_fps(self) -> float | None: return self._throughput(self.fans_processed) @property def throughput_rps(self) -> float | None: return self._throughput(self.total_requests) def _throughput(self, value: int) -> float | None: if not self.started_at or value <= 0: return None end = self.finished_at or timezone.now() elapsed_s = (end - self.started_at).total_seconds() if elapsed_s <= 0: return None return round(value / elapsed_s, 1)