"""Async store for concurrency samples + minute/hour aggregation.""" from __future__ import annotations import logging from dataclasses import dataclass from datetime import UTC, datetime, timedelta from sqlalchemy import delete, select from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker from deerflow.persistence.concurrency.model import ConcurrencySampleRow logger = logging.getLogger(__name__) # How far back the aggregation will fetch raw rows before bucketing. A safety # cap so a degenerate range never tries to load the whole table into memory. _MAX_ROWS = 60_000 @dataclass class ConcurrencyBucket: """One minute/hour bucket of the concurrency series.""" bucket: str # ISO start of the bucket peak: int # max concurrent runs observed in the bucket avg: float # mean concurrent runs in the bucket samples: int # how many raw samples fell into the bucket def _bucket_start(dt: datetime, granularity: str) -> datetime: """Truncate a timestamp to the start of its minute/hour bucket.""" if granularity == "hour": return dt.replace(minute=0, second=0, microsecond=0) return dt.replace(second=0, microsecond=0) class ConcurrencySampleStore: """Persists periodic concurrency snapshots and aggregates them.""" def __init__(self, session_factory: async_sessionmaker[AsyncSession]) -> None: self._sf = session_factory async def record(self, active_count: int, *, sampled_at: datetime | None = None) -> None: """Insert one sample. Best-effort — never raises into the caller.""" try: async with self._sf() as session: session.add( ConcurrencySampleRow( active_count=max(0, int(active_count)), sampled_at=sampled_at or datetime.now(UTC), ) ) await session.commit() except Exception: # pragma: no cover - sampling must never break the app logger.warning("Failed to record concurrency sample", exc_info=True) async def aggregate( self, *, since: datetime, until: datetime, granularity: str = "minute", ) -> list[ConcurrencyBucket]: """Return per-minute / per-hour buckets of peak + average concurrency. Bucketing is done in Python (not SQL ``date_trunc``/``strftime``) so the logic is identical across the sqlite/mysql/postgres backends. """ gran = "hour" if granularity == "hour" else "minute" try: async with self._sf() as session: rows = ( await session.execute( select(ConcurrencySampleRow.sampled_at, ConcurrencySampleRow.active_count) .where( ConcurrencySampleRow.sampled_at >= since, ConcurrencySampleRow.sampled_at <= until, ) .order_by(ConcurrencySampleRow.sampled_at) .limit(_MAX_ROWS) ) ).all() except Exception: logger.warning("Failed to aggregate concurrency samples", exc_info=True) return [] buckets: dict[datetime, dict[str, float]] = {} for sampled_at, active_count in rows: if sampled_at is None: continue if sampled_at.tzinfo is None: sampled_at = sampled_at.replace(tzinfo=UTC) key = _bucket_start(sampled_at, gran) value = int(active_count or 0) entry = buckets.get(key) if entry is None: buckets[key] = {"peak": value, "sum": value, "count": 1} else: entry["peak"] = max(entry["peak"], value) entry["sum"] += value entry["count"] += 1 result: list[ConcurrencyBucket] = [] for key in sorted(buckets): entry = buckets[key] count = max(1, int(entry["count"])) result.append( ConcurrencyBucket( bucket=key.isoformat(), peak=int(entry["peak"]), avg=round(entry["sum"] / count, 2), samples=count, ) ) return result async def purge_older_than(self, cutoff: datetime) -> int: """Delete samples older than ``cutoff``. Returns rows removed.""" try: async with self._sf() as session: result = await session.execute( delete(ConcurrencySampleRow).where(ConcurrencySampleRow.sampled_at < cutoff) ) await session.commit() return int(result.rowcount or 0) except Exception: # pragma: no cover logger.warning("Failed to purge concurrency samples", exc_info=True) return 0 def make_concurrency_sample_store( session_factory: async_sessionmaker[AsyncSession] | None, ) -> ConcurrencySampleStore | None: """Factory mirroring the other ``make_*_store`` helpers. Returns ``None`` on the memory backend (no SQL session factory), in which case the monitor simply reports a live snapshot with no historical chart. """ if session_factory is None: return None return ConcurrencySampleStore(session_factory)