143 lines
5.3 KiB
Python
143 lines
5.3 KiB
Python
"""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)
|