Record per-case retrieval diagnostics in evaluations

Count tool traffic from the message history, where refused and repeated
calls stay visible, unlike state.searches which is keyed by query.
This commit is contained in:
Yiorgis Gozadinos 2026-07-29 15:30:26 +03:00
parent 5752f61f2c
commit a528ab912f
No known key found for this signature in database
4 changed files with 104 additions and 2 deletions

View file

@ -1,6 +1,10 @@
# Changelog
## [Unreleased]
### Added
- `evaluations run` records `cited_chunk_ids`, `searched_uris`, `n_searches`, `n_search_calls`, `n_rejected_calls`, `n_executions`, `n_requests` and `budget_spent` as eval attributes alongside `cited_uris`.
### Changed
- A capability whose search or code-execution budget is spent says so in its instructions on every following request, naming the exhausted tools.

View file

@ -434,6 +434,14 @@ async def run_qa_benchmark(
capability_model=resolved_capability_model,
)
set_eval_attribute("cited_uris", result.cited_uris)
set_eval_attribute("cited_chunk_ids", result.cited_chunk_ids)
set_eval_attribute("searched_uris", result.searched_uris)
set_eval_attribute("n_searches", result.n_searches)
set_eval_attribute("n_search_calls", result.n_search_calls)
set_eval_attribute("n_rejected_calls", result.n_rejected_calls)
set_eval_attribute("n_executions", result.n_executions)
set_eval_attribute("n_requests", result.n_requests)
set_eval_attribute("budget_spent", result.budget_spent)
return result.answer
report = await _evaluate(answer_question)

View file

@ -4,6 +4,12 @@ from pathlib import Path
from typing import Any, Protocol, cast
from pydantic_ai import Agent
from pydantic_ai.messages import (
ModelMessage,
ModelResponse,
ToolCallPart,
ToolReturnPart,
)
from pydantic_ai.models import Model
from haiku.rag.capabilities import RAGCapabilityBase
@ -29,6 +35,41 @@ class CapabilityRunResult:
searched_uris: list[str] = field(default_factory=list)
n_searches: int = 0
n_executions: int = 0
n_search_calls: int = 0
n_rejected_calls: int = 0
n_requests: int = 0
budget_spent: bool = False
def _count_tool_traffic(
messages: list[ModelMessage], namespace: str
) -> tuple[int, int, int]:
"""Count search calls, rejected calls and model requests in a run.
``state.searches`` is keyed by query, so it collapses repeated queries and
never records a call the capability refused. Counting the message history
instead gives the real number of attempts, which is what shows whether a
case ran out of budget.
"""
search_tool = f"{namespace}_search"
search_calls = 0
rejected = 0
requests = 0
for message in messages:
if isinstance(message, ModelResponse):
requests += 1
search_calls += sum(
1
for part in message.parts
if isinstance(part, ToolCallPart) and part.tool_name == search_tool
)
continue
rejected += sum(
1
for part in message.parts
if isinstance(part, ToolReturnPart) and part.outcome == "failed"
)
return search_calls, rejected, requests
@dataclass
@ -100,6 +141,10 @@ async def run_capability_question(
executions = getattr(state, "executions", None)
n_executions = len(executions) if executions is not None else 0
n_search_calls, n_rejected_calls, n_requests = _count_tool_traffic(
agent_result.all_messages(), capability.state_namespace
)
return CapabilityRunResult(
answer=agent_result.output,
cited_uris=cited_uris,
@ -107,4 +152,8 @@ async def run_capability_question(
searched_uris=searched_uris,
n_searches=len(typed.searches),
n_executions=n_executions,
n_search_calls=n_search_calls,
n_rejected_calls=n_rejected_calls,
n_requests=n_requests,
budget_spent=n_rejected_calls > 0,
)

View file

@ -2,14 +2,55 @@ from types import SimpleNamespace
from unittest.mock import AsyncMock, patch
import pytest
from pydantic_ai.messages import (
ModelRequest,
ModelResponse,
TextPart,
ToolCallPart,
ToolReturnPart,
UserPromptPart,
)
from pydantic_ai.models.test import TestModel
from evaluations.capability_runner import run_capability_question
from evaluations.capability_runner import _count_tool_traffic, run_capability_question
from haiku.rag.capabilities.analysis import create_capability as create_analysis
from haiku.rag.capabilities.rag import create_capability as create_rag
from haiku.rag.config.models import AppConfig
def test_count_tool_traffic_counts_attempts_not_distinct_queries():
"""Rejected and repeated calls both count; `state.searches` hides them."""
messages = [
ModelRequest(parts=[UserPromptPart(content="q")]),
ModelResponse(
parts=[
ToolCallPart("analysis_search", {"query": "same"}),
ToolCallPart("analysis_search", {"query": "same"}),
]
),
ModelRequest(
parts=[
ToolReturnPart(
tool_name="analysis_search", content="results", tool_call_id="1"
),
ToolReturnPart(
tool_name="analysis_search",
content="Search limit reached.",
tool_call_id="2",
outcome="failed",
),
]
),
ModelResponse(parts=[TextPart("done")]),
]
search_calls, rejected, requests = _count_tool_traffic(messages, "analysis")
assert search_calls == 2
assert rejected == 1
assert requests == 2
async def test_runs_rag_capability_without_legacy_capability_layer(tmp_path):
result = await run_capability_question(
create_rag,
@ -49,7 +90,7 @@ async def test_analysis_capability_applies_request_limit(tmp_path, override, exp
with patch(
"evaluations.capability_runner.Agent.run", new_callable=AsyncMock
) as run:
run.return_value = SimpleNamespace(output="done")
run.return_value = SimpleNamespace(output="done", all_messages=lambda: [])
await run_capability_question(
lambda **_kwargs: capability,