248 lines
10 KiB
Python
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)
|
|
]
|