113 lines
4.3 KiB
Python
113 lines
4.3 KiB
Python
"""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
|