diff --git a/agent/insights.py b/agent/insights.py index 70907b4f3d57c..6ebd1abc90efe 100644 --- a/agent/insights.py +++ b/agent/insights.py @@ -212,39 +212,39 @@ def _get_tool_usage(self, cutoff: float, source: str = None) -> List[Dict]: 2. tool_calls JSON on 'assistant' role messages (covers CLI where tool_name is not populated on tool responses) """ - tool_counts = Counter() + tool_counts_by_session = Counter() # Source 1: explicit tool_name on tool response messages if source: cursor = self._conn.execute( - """SELECT m.tool_name, COUNT(*) as count + """SELECT m.session_id, m.tool_name, COUNT(*) as count FROM messages m JOIN sessions s ON s.id = m.session_id WHERE s.started_at >= ? AND s.source = ? AND m.role = 'tool' AND m.tool_name IS NOT NULL - GROUP BY m.tool_name + GROUP BY m.session_id, m.tool_name ORDER BY count DESC""", (cutoff, source), ) else: cursor = self._conn.execute( - """SELECT m.tool_name, COUNT(*) as count + """SELECT m.session_id, m.tool_name, COUNT(*) as count FROM messages m JOIN sessions s ON s.id = m.session_id WHERE s.started_at >= ? AND m.role = 'tool' AND m.tool_name IS NOT NULL - GROUP BY m.tool_name + GROUP BY m.session_id, m.tool_name ORDER BY count DESC""", (cutoff,), ) for row in cursor.fetchall(): - tool_counts[row["tool_name"]] += row["count"] + tool_counts_by_session[(row["session_id"], row["tool_name"])] += row["count"] # Source 2: extract from tool_calls JSON on assistant messages # (covers CLI sessions where tool_name is NULL on tool responses) if source: cursor2 = self._conn.execute( - """SELECT m.tool_calls + """SELECT m.session_id, m.tool_calls FROM messages m JOIN sessions s ON s.id = m.session_id WHERE s.started_at >= ? AND s.source = ? @@ -253,7 +253,7 @@ def _get_tool_usage(self, cutoff: float, source: str = None) -> List[Dict]: ) else: cursor2 = self._conn.execute( - """SELECT m.tool_calls + """SELECT m.session_id, m.tool_calls FROM messages m JOIN sessions s ON s.id = m.session_id WHERE s.started_at >= ? @@ -261,7 +261,7 @@ def _get_tool_usage(self, cutoff: float, source: str = None) -> List[Dict]: (cutoff,), ) - tool_calls_counts = Counter() + tool_calls_counts_by_session = Counter() for row in cursor2.fetchall(): try: calls = row["tool_calls"] @@ -272,23 +272,24 @@ def _get_tool_usage(self, cutoff: float, source: str = None) -> List[Dict]: func = call.get("function", {}) if isinstance(call, dict) else {} name = func.get("name") if name: - tool_calls_counts[name] += 1 + tool_calls_counts_by_session[(row["session_id"], name)] += 1 except (json.JSONDecodeError, TypeError, AttributeError): continue - # Merge: prefer tool_name source, supplement with tool_calls source - # for tools not already counted - if not tool_counts and tool_calls_counts: - # No tool_name data at all — use tool_calls exclusively - tool_counts = tool_calls_counts - elif tool_counts and tool_calls_counts: - # Both sources have data — use whichever has the higher count per tool - # (they may overlap, so take the max to avoid double-counting) - all_tools = set(tool_counts) | set(tool_calls_counts) - merged = Counter() - for tool in all_tools: - merged[tool] = max(tool_counts.get(tool, 0), tool_calls_counts.get(tool, 0)) - tool_counts = merged + # Merge per session/tool so mixed datasets across different sessions add + # together, while duplicated representations inside the same session + # still collapse to the higher count. + merged_by_session = Counter() + all_session_tools = set(tool_counts_by_session) | set(tool_calls_counts_by_session) + for session_tool in all_session_tools: + merged_by_session[session_tool] = max( + tool_counts_by_session.get(session_tool, 0), + tool_calls_counts_by_session.get(session_tool, 0), + ) + + tool_counts = Counter() + for (_, tool_name), count in merged_by_session.items(): + tool_counts[tool_name] += count # Convert to the expected format return [ diff --git a/tests/agent/test_insights.py b/tests/agent/test_insights.py index 2740daf0962fa..86a295a17a076 100644 --- a/tests/agent/test_insights.py +++ b/tests/agent/test_insights.py @@ -630,6 +630,30 @@ def test_tool_usage_from_tool_calls_json(self, db): sf = next(t for t in tools if t["tool"] == "search_files") assert sf["count"] == 2 + def test_tool_usage_sums_disjoint_sessions_across_sources(self, db): + """tool_name and tool_calls sources should add across different sessions.""" + db.create_session(session_id="gateway", source="gateway", model="test") + db.append_message("gateway", role="tool", content="results", tool_name="search_files") + + db.create_session(session_id="cli", source="cli", model="test") + db.append_message( + "cli", + role="assistant", + content="Let me search", + tool_calls=[{ + "id": "call_1", + "type": "function", + "function": {"name": "search_files", "arguments": "{}"}, + }], + ) + db._conn.commit() + + engine = InsightsEngine(db) + report = engine.generate(days=30) + + search_files = next(t for t in report["tools"] if t["tool"] == "search_files") + assert search_files["count"] == 2 + def test_overview_pricing_sets_are_lists(self, db): """models_with/without_pricing should be JSON-serializable lists.""" import json as _json