"""SQLAlchemy-backed LLM metrics repository.""" from __future__ import annotations import logging from collections.abc import AsyncIterator from datetime import UTC, datetime from typing import Any from sqlalchemy import desc, func, select from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker from deerflow.persistence.llm_metrics.base import LlmMetricRecord, LlmMetricsQuery, LlmMetricsStore from deerflow.persistence.llm_metrics.model import LlmCallMetricRow logger = logging.getLogger(__name__) def _row_to_dict(row: LlmCallMetricRow) -> dict[str, Any]: return row.to_dict() def _normalize_created_at(value: datetime) -> datetime: return value if value.tzinfo else value.replace(tzinfo=UTC) class SqlLlmMetricsStore(LlmMetricsStore): """SQLite/Postgres-backed implementation; writes are intentionally fire-and-forget.""" def __init__(self, session_factory: async_sessionmaker[AsyncSession]) -> None: self._sf = session_factory async def record(self, metric: LlmMetricRecord) -> None: row = LlmCallMetricRow( id=metric.id, created_at=_normalize_created_at(metric.created_at), user_id=metric.user_id, thread_id=metric.thread_id, run_id=metric.run_id, agent_name=metric.agent_name, model_name=metric.model_name, duration_ms=metric.duration_ms, input_tokens=metric.input_tokens, output_tokens=metric.output_tokens, total_tokens=metric.total_tokens, tokens_per_sec=metric.tokens_per_sec, status=metric.status, error_type=metric.error_type, error_message=metric.error_message, ) try: async with self._sf() as session: session.add(row) await session.commit() except Exception: # Metrics are best-effort. Never let a persistence failure surface # to the agent runtime — log and move on. logger.warning("Failed to persist llm_call_metrics row %s", metric.id, exc_info=True) def _apply_filters(self, stmt, query: LlmMetricsQuery): if query.user_ids: stmt = stmt.where(LlmCallMetricRow.user_id.in_(query.user_ids)) elif query.user_id: stmt = stmt.where(LlmCallMetricRow.user_id == query.user_id) if query.model_name: stmt = stmt.where(LlmCallMetricRow.model_name == query.model_name) if query.status: stmt = stmt.where(LlmCallMetricRow.status == query.status) if query.since is not None: stmt = stmt.where(LlmCallMetricRow.created_at >= _normalize_created_at(query.since)) if query.until is not None: stmt = stmt.where(LlmCallMetricRow.created_at <= _normalize_created_at(query.until)) return stmt async def list(self, query: LlmMetricsQuery) -> list[dict[str, Any]]: stmt = select(LlmCallMetricRow) stmt = self._apply_filters(stmt, query) stmt = stmt.order_by(desc(LlmCallMetricRow.created_at)).limit(max(1, min(query.limit, 1000))).offset(max(0, query.offset)) async with self._sf() as session: result = await session.execute(stmt) return [_row_to_dict(row) for row in result.scalars()] async def count(self, query: LlmMetricsQuery) -> int: stmt = select(func.count()).select_from(LlmCallMetricRow) stmt = self._apply_filters(stmt, query) async with self._sf() as session: return (await session.scalar(stmt)) or 0 async def iter_all(self, query: LlmMetricsQuery) -> AsyncIterator[dict[str, Any]]: # Page through the result so the Excel export stays bounded in memory # for large date windows. page_size = 500 offset = 0 while True: page_query = LlmMetricsQuery( user_id=query.user_id, user_ids=query.user_ids, model_name=query.model_name, status=query.status, since=query.since, until=query.until, limit=page_size, offset=offset, ) page = await self.list(page_query) if not page: return for item in page: yield item if len(page) < page_size: return offset += page_size