The request-limit notice said only the cite tool remained available, but chat registers rag and analysis in one agent, so exhausting analysis claimed rag_search was gone too. Scoped to the capability's own tools. The cite window was counted over every model request once loaded, so turns spent on another capability expired it before the model was ever placed where citing was the obvious move. Count only requests whose preceding response called one of this capability's tools; engagement is also the only thing that can loop, which is all the bound guards against. Also: _count_tool_traffic returns a named tuple rather than four bare ints, and counts failures only for this capability's tools, so host-tool retries and output-validation retries no longer read as its failures.
159 lines
5 KiB
Python
159 lines
5 KiB
Python
from types import SimpleNamespace
|
|
from unittest.mock import AsyncMock, patch
|
|
|
|
import pytest
|
|
from pydantic_ai.messages import (
|
|
ModelRequest,
|
|
ModelResponse,
|
|
RetryPromptPart,
|
|
TextPart,
|
|
ToolCallPart,
|
|
ToolReturnPart,
|
|
UserPromptPart,
|
|
)
|
|
from pydantic_ai.models.test import TestModel
|
|
|
|
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
|
|
|
|
ANALYSIS_TOOLS = frozenset(
|
|
{"analysis_search", "analysis_execute_code", "analysis_cite"}
|
|
)
|
|
|
|
|
|
def test_count_tool_traffic_sees_a_rejected_cite_call():
|
|
"""`_cite` rejects with ModelRetry, which is not a failed ToolReturnPart."""
|
|
messages = [
|
|
ModelRequest(parts=[UserPromptPart(content="q")]),
|
|
ModelResponse(parts=[ToolCallPart("analysis_cite", {"chunk_ids": []})]),
|
|
ModelRequest(
|
|
parts=[
|
|
RetryPromptPart(
|
|
tool_name="analysis_cite",
|
|
content="No citations registered: chunk_ids was empty.",
|
|
tool_call_id="1",
|
|
)
|
|
]
|
|
),
|
|
ModelResponse(parts=[TextPart("done")]),
|
|
]
|
|
|
|
traffic = _count_tool_traffic(messages, "analysis", ANALYSIS_TOOLS)
|
|
|
|
assert traffic.n_failed_tools == 1
|
|
assert traffic.n_rejected_searches == 0
|
|
|
|
|
|
def test_count_tool_traffic_separates_search_rejections_from_code_errors():
|
|
"""A crash in model-written Python must not read as budget exhaustion."""
|
|
messages = [
|
|
ModelRequest(parts=[UserPromptPart(content="q")]),
|
|
ModelResponse(parts=[ToolCallPart("analysis_execute_code", {"code": "1/0"})]),
|
|
ModelRequest(
|
|
parts=[
|
|
ToolReturnPart(
|
|
tool_name="analysis_execute_code",
|
|
content="ZeroDivisionError",
|
|
tool_call_id="1",
|
|
outcome="failed",
|
|
)
|
|
]
|
|
),
|
|
ModelResponse(parts=[TextPart("done")]),
|
|
]
|
|
|
|
traffic = _count_tool_traffic(messages, "analysis", ANALYSIS_TOOLS)
|
|
|
|
assert traffic.n_search_calls == 0
|
|
assert traffic.n_rejected_searches == 0
|
|
assert traffic.n_failed_tools == 1
|
|
assert traffic.n_requests == 2
|
|
|
|
|
|
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")]),
|
|
]
|
|
|
|
traffic = _count_tool_traffic(messages, "analysis", ANALYSIS_TOOLS)
|
|
|
|
assert traffic.n_search_calls == 2
|
|
assert traffic.n_rejected_searches == 1
|
|
assert traffic.n_requests == 2
|
|
|
|
|
|
async def test_runs_rag_capability_without_legacy_capability_layer(tmp_path):
|
|
result = await run_capability_question(
|
|
create_rag,
|
|
tmp_path / "rag.lancedb",
|
|
AppConfig(),
|
|
"hello",
|
|
TestModel(call_tools=[]),
|
|
document_filter="uri = 'manual.pdf'",
|
|
)
|
|
|
|
assert result.answer == "success (no tool calls)"
|
|
assert result.cited_uris == []
|
|
assert result.n_searches == 0
|
|
|
|
|
|
async def test_runs_analysis_capability_without_legacy_capability_layer(tmp_path):
|
|
result = await run_capability_question(
|
|
create_analysis,
|
|
tmp_path / "rag.lancedb",
|
|
AppConfig(),
|
|
"hello",
|
|
TestModel(call_tools=[]),
|
|
request_limit=5,
|
|
)
|
|
|
|
assert result.answer == "success (no tool calls)"
|
|
assert result.n_executions == 0
|
|
|
|
|
|
@pytest.mark.parametrize(("override", "expected"), [(None, 30), (5, 5)])
|
|
async def test_analysis_capability_applies_request_limit(tmp_path, override, expected):
|
|
capability = create_analysis(
|
|
db_path=tmp_path / "rag.lancedb",
|
|
config=AppConfig(),
|
|
defer_loading=False,
|
|
)
|
|
with patch(
|
|
"evaluations.capability_runner.Agent.run", new_callable=AsyncMock
|
|
) as run:
|
|
run.return_value = SimpleNamespace(output="done", all_messages=lambda: [])
|
|
|
|
await run_capability_question(
|
|
lambda **_kwargs: capability,
|
|
tmp_path / "rag.lancedb",
|
|
AppConfig(),
|
|
"hello",
|
|
TestModel(call_tools=[]),
|
|
request_limit=override,
|
|
)
|
|
|
|
assert capability.request_limit == expected
|
|
assert "usage_limits" not in run.call_args.kwargs
|