"""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) ]