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

248 lines
10 KiB
Python

"""SQLAlchemy-backed tool metrics repository."""
from __future__ import annotations
import asyncio
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.tool_metrics.base import ToolMetricRecord, ToolMetricsQuery, ToolMetricsStore
from deerflow.persistence.tool_metrics.model import ToolCallMetricRow
logger = logging.getLogger(__name__)
# Separator joining every failed attempt's message inside one merged skill row.
_ERROR_SEP = "\n--- \n"
# Hard cap on the accumulated ``error_message`` of a merged row.
_MERGED_ERROR_LIMIT = 2000
# A single skill invocation often spans several tool calls (read SKILL.md, run
# a script, ...), and the model may retry a skill within the same run. Those
# fire-and-forget writes can interleave on the event loop, so the read → merge →
# write of a skill row is guarded by a lock. ``asyncio.Lock`` binds to the loop
# it is used on; the lock is therefore (re)created whenever the running loop
# changes (the Gateway keeps one persistent loop, while the ``asyncio.run``
# fallback path spins up a throwaway loop per call — already serialized).
_merge_lock: asyncio.Lock | None = None
_merge_lock_loop: object = None
def _skill_merge_lock() -> asyncio.Lock:
global _merge_lock, _merge_lock_loop
loop = asyncio.get_running_loop()
if _merge_lock is None or _merge_lock_loop is not loop:
_merge_lock = asyncio.Lock()
_merge_lock_loop = loop
return _merge_lock
def _normalize_created_at(value: datetime) -> datetime:
return value if value.tzinfo else value.replace(tzinfo=UTC)
def _merge_error_messages(existing: str | None, new: str | None) -> str | None:
"""Append a failed attempt's message to the accumulated error text.
All failures are kept ("都统计下来"), separated and capped at
:data:`_MERGED_ERROR_LIMIT` characters.
"""
new = (new or "").strip()
if not new:
return existing
existing = (existing or "").strip()
if not existing:
return new
if new in existing:
# Identical retry failure — no point storing the same line twice.
return existing
combined = existing + _ERROR_SEP + new
if len(combined) > _MERGED_ERROR_LIMIT:
combined = combined[: _MERGED_ERROR_LIMIT - 3] + "..."
return combined
def _merge_skill_row(existing: ToolCallMetricRow, metric: ToolMetricRecord) -> None:
"""Fold another same-run, same-skill call into an existing metric row.
Outcome rules:
- any success in the run -> the row becomes a single ``success`` and earlier
failures are dropped ("先报错、最后成功 -> 记一次成功");
- all calls failed -> a single ``error`` row whose ``error_message``
accumulates every attempt's message ("一直报错 -> 记一次,报错信息都统计下来").
"""
existing.duration_ms = (existing.duration_ms or 0) + (metric.duration_ms or 0)
if metric.status == "success":
if existing.status != "success":
existing.status = "success"
existing.error_category = None
existing.error_message = None
return
# metric.status == "error"
if existing.status == "success":
# The skill already succeeded once this run — keep that outcome.
return
existing.status = "error"
if metric.error_category:
existing.error_category = metric.error_category
existing.error_message = _merge_error_messages(existing.error_message, metric.error_message)
def _row_from_metric(metric: ToolMetricRecord) -> ToolCallMetricRow:
return ToolCallMetricRow(
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,
tool_name=metric.tool_name,
skill_name=metric.skill_name,
duration_ms=metric.duration_ms,
status=metric.status,
error_category=metric.error_category,
error_message=metric.error_message,
)
class SqlToolMetricsStore(ToolMetricsStore):
"""SQLite/Postgres/MySQL-backed implementation; writes are fire-and-forget."""
def __init__(self, session_factory: async_sessionmaker[AsyncSession]) -> None:
self._sf = session_factory
async def record(self, metric: ToolMetricRecord) -> None:
# Skill calls are merged per run: the model may touch one skill across
# several tool calls / retries, and the admin wants a single row per
# (run, skill) instead of one per call. Non-skill tool calls, and skill
# calls with no run_id, fall back to a plain insert.
if metric.skill_name and metric.run_id:
await self._record_skill_merged(metric)
return
try:
async with self._sf() as session:
session.add(_row_from_metric(metric))
await session.commit()
except Exception:
# Metrics are best-effort; never surface a write failure to the runtime.
logger.warning("Failed to persist tool_call_metrics row %s", metric.id, exc_info=True)
async def _record_skill_merged(self, metric: ToolMetricRecord) -> None:
"""Upsert a per-(run, skill) row, folding repeated calls into one."""
try:
async with _skill_merge_lock():
async with self._sf() as session:
existing = await session.scalar(
select(ToolCallMetricRow)
.where(
ToolCallMetricRow.run_id == metric.run_id,
ToolCallMetricRow.skill_name == metric.skill_name,
)
.order_by(ToolCallMetricRow.created_at)
.limit(1)
)
if existing is None:
session.add(_row_from_metric(metric))
else:
_merge_skill_row(existing, metric)
await session.commit()
except Exception:
logger.warning(
"Failed to merge tool_call_metrics for skill %s (run %s)",
metric.skill_name,
metric.run_id,
exc_info=True,
)
def _apply_filters(self, stmt, query: ToolMetricsQuery):
if query.user_id:
stmt = stmt.where(ToolCallMetricRow.user_id == query.user_id)
if query.tool_name:
stmt = stmt.where(ToolCallMetricRow.tool_name == query.tool_name)
if query.skill_only:
stmt = stmt.where(ToolCallMetricRow.skill_name.is_not(None))
if query.skill_name:
stmt = stmt.where(ToolCallMetricRow.skill_name == query.skill_name)
if query.status:
stmt = stmt.where(ToolCallMetricRow.status == query.status)
if query.error_category:
stmt = stmt.where(ToolCallMetricRow.error_category == query.error_category)
if query.since is not None:
stmt = stmt.where(ToolCallMetricRow.created_at >= _normalize_created_at(query.since))
if query.until is not None:
stmt = stmt.where(ToolCallMetricRow.created_at <= _normalize_created_at(query.until))
return stmt
async def list(self, query: ToolMetricsQuery) -> list[dict[str, Any]]:
stmt = select(ToolCallMetricRow)
stmt = self._apply_filters(stmt, query)
stmt = stmt.order_by(desc(ToolCallMetricRow.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() for row in result.scalars()]
async def count(self, query: ToolMetricsQuery) -> int:
stmt = select(func.count()).select_from(ToolCallMetricRow)
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: ToolMetricsQuery) -> AsyncIterator[dict[str, Any]]:
page_size = 500
offset = 0
while True:
page = await self.list(
ToolMetricsQuery(
user_id=query.user_id,
tool_name=query.tool_name,
skill_only=query.skill_only,
skill_name=query.skill_name,
status=query.status,
error_category=query.error_category,
since=query.since,
until=query.until,
limit=page_size,
offset=offset,
)
)
if not page:
return
for item in page:
yield item
if len(page) < page_size:
return
offset += page_size
async def skill_breakdown(self, query: ToolMetricsQuery) -> list[dict[str, Any]]:
# Count success/error per skill. Force the skill-only filter so rows
# without a skill_name (plain tool calls) never leak into the result.
scoped = ToolMetricsQuery(
user_id=query.user_id,
skill_only=True,
skill_name=query.skill_name,
since=query.since,
until=query.until,
)
stmt = select(
ToolCallMetricRow.skill_name,
ToolCallMetricRow.status,
func.count().label("n"),
)
stmt = self._apply_filters(stmt, scoped).group_by(
ToolCallMetricRow.skill_name, ToolCallMetricRow.status
)
agg: dict[str, dict[str, int]] = {}
async with self._sf() as session:
for skill_name, status, n in (await session.execute(stmt)).all():
bucket = agg.setdefault(skill_name, {"total": 0, "success": 0, "error": 0})
bucket["total"] += n
bucket["success" if status == "success" else "error"] += n
return [
{"skill_name": name, **counts}
for name, counts in sorted(agg.items(), key=lambda kv: kv[1]["total"], reverse=True)
]