Thread per-case source scope through the capability runner

The runner passed document_filter onto capability state but not sources,
so an eval case could not be scoped to named databases. Scoping is what
the multi-database acceptance dataset has to assert, both that a scoped
question honours its databases and that an empty scope refuses rather
than answering without evidence.

sources=[] covers no database while None covers every one the client
covers, so the empty list is threaded on an is-not-None check and a test
pins the two apart.
This commit is contained in:
Yiorgis Gozadinos 2026-08-26 14:56:46 +03:00
parent 8ed24d0e24
commit ba760b76c2
No known key found for this signature in database
2 changed files with 86 additions and 2 deletions

View file

@ -130,6 +130,7 @@ def _prepare_agent(
document_filter: str | None,
request_limit: int | None,
compaction: bool = False,
sources: list[str] | None = None,
) -> tuple[RAGCapabilityBase[Any], _EvalDeps, Agent[_EvalDeps, str]]:
capability = capability_factory(
db_path=db_path,
@ -141,6 +142,8 @@ def _prepare_agent(
state = capability.state_type()
if document_filter is not None:
state.document_filter = document_filter
if sources is not None:
state.sources = sources
capabilities: list[AbstractCapability] = [capability]
if compaction:
@ -169,6 +172,7 @@ async def run_capability_question(
document_filter: str | None = None,
request_limit: int | None = None,
message_history: list[ModelMessage] | None = None,
sources: list[str] | None = None,
) -> CapabilityRunResult:
"""Run a single question through a capability and return answer + retrieval data.
@ -177,8 +181,11 @@ async def run_capability_question(
are extracted from the state for downstream eval scoring.
The capability must produce a state with RAG-capability-shaped fields (citation
index, searches, optional document filter) i.e. ``RAGState`` or
``AnalysisState`` from ``haiku.rag.capabilities``.
index, searches, optional document filter and source scope) i.e. ``RAGState``
or ``AnalysisState`` from ``haiku.rag.capabilities``.
``sources`` scopes the question to named databases; ``[]`` covers none and
``None`` covers everything the client covers.
"""
capability, deps, agent = _prepare_agent(
capability_factory,
@ -187,6 +194,7 @@ async def run_capability_question(
capability_model,
document_filter,
request_limit,
sources=sources,
)
agent_result = await agent.run(question, deps=deps, message_history=message_history)
traffic = _count_tool_traffic(
@ -205,6 +213,7 @@ async def run_capability_conversation(
capability_model: str | Model,
document_filter: str | None = None,
compaction: bool = False,
sources: list[str] | None = None,
) -> list[CapabilityRunResult]:
"""Run a conversation's user turns sequentially through one capability.
@ -223,6 +232,7 @@ async def run_capability_conversation(
document_filter=document_filter,
request_limit=None,
compaction=compaction,
sources=sources,
)
history: list[ModelMessage] | None = None
results: list[CapabilityRunResult] = []

View file

@ -345,6 +345,80 @@ async def test_conversation_applies_document_filter(tmp_path):
assert deps_seen[0].state["rag"]["document_filter"] == "uri = 'manual.pdf'"
async def test_question_applies_sources(tmp_path):
"""Per-case `sources` must reach the capability state, so a scoped question
only ever searches the databases it named."""
deps_seen = []
async def _run(question, deps=None, message_history=None):
deps_seen.append(deps)
return SimpleNamespace(output="a", new_messages=lambda: [])
with patch("evaluations.capability_runner.Agent.run", side_effect=_run):
await run_capability_question(
create_rag,
tmp_path / "rag.lancedb",
AppConfig(),
"q",
TestModel(call_tools=[]),
sources=["northern"],
)
assert deps_seen[0].state["rag"]["sources"] == ["northern"]
async def test_question_distinguishes_empty_sources_from_none(tmp_path):
"""`sources=[]` covers nothing and `sources=None` covers everything, so the
empty list must survive to the state rather than being treated as unset."""
deps_seen = []
async def _run(question, deps=None, message_history=None):
deps_seen.append(deps)
return SimpleNamespace(output="a", new_messages=lambda: [])
with patch("evaluations.capability_runner.Agent.run", side_effect=_run):
for sources in ([], None):
await run_capability_question(
create_rag,
tmp_path / "rag.lancedb",
AppConfig(),
"q",
TestModel(call_tools=[]),
sources=sources,
)
assert deps_seen[0].state["rag"]["sources"] == []
assert deps_seen[1].state["rag"]["sources"] is None
async def test_conversation_applies_sources(tmp_path):
"""Scope must hold for every turn of a conversation, not just the first."""
from evaluations.capability_runner import run_capability_conversation
deps_seen = []
async def _run(question, deps=None, message_history=None):
deps_seen.append(deps)
return SimpleNamespace(
output="a", all_messages=lambda: [], new_messages=lambda: []
)
with patch("evaluations.capability_runner.Agent.run", side_effect=_run):
await run_capability_conversation(
create_rag,
tmp_path / "rag.lancedb",
AppConfig(),
["q1", "q2"],
TestModel(call_tools=[]),
sources=["southern"],
)
assert [d.state["rag"]["sources"] for d in deps_seen] == [
["southern"],
["southern"],
]
async def test_conversation_carries_one_state_dict_across_turns(tmp_path):
"""Capabilities read and write state through the deps dict; carrying the
same dict across turns is what lets compaction see earlier questions'