deerflow-code/offline-backend-20260512/backend/packages/harness/deerflow/persistence/llm_metrics/sql.py
2026-09-07 18:24:55 +08:00

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