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:
parent
8ed24d0e24
commit
ba760b76c2
2 changed files with 86 additions and 2 deletions
|
|
@ -130,6 +130,7 @@ def _prepare_agent(
|
||||||
document_filter: str | None,
|
document_filter: str | None,
|
||||||
request_limit: int | None,
|
request_limit: int | None,
|
||||||
compaction: bool = False,
|
compaction: bool = False,
|
||||||
|
sources: list[str] | None = None,
|
||||||
) -> tuple[RAGCapabilityBase[Any], _EvalDeps, Agent[_EvalDeps, str]]:
|
) -> tuple[RAGCapabilityBase[Any], _EvalDeps, Agent[_EvalDeps, str]]:
|
||||||
capability = capability_factory(
|
capability = capability_factory(
|
||||||
db_path=db_path,
|
db_path=db_path,
|
||||||
|
|
@ -141,6 +142,8 @@ def _prepare_agent(
|
||||||
state = capability.state_type()
|
state = capability.state_type()
|
||||||
if document_filter is not None:
|
if document_filter is not None:
|
||||||
state.document_filter = document_filter
|
state.document_filter = document_filter
|
||||||
|
if sources is not None:
|
||||||
|
state.sources = sources
|
||||||
|
|
||||||
capabilities: list[AbstractCapability] = [capability]
|
capabilities: list[AbstractCapability] = [capability]
|
||||||
if compaction:
|
if compaction:
|
||||||
|
|
@ -169,6 +172,7 @@ async def run_capability_question(
|
||||||
document_filter: str | None = None,
|
document_filter: str | None = None,
|
||||||
request_limit: int | None = None,
|
request_limit: int | None = None,
|
||||||
message_history: list[ModelMessage] | None = None,
|
message_history: list[ModelMessage] | None = None,
|
||||||
|
sources: list[str] | None = None,
|
||||||
) -> CapabilityRunResult:
|
) -> CapabilityRunResult:
|
||||||
"""Run a single question through a capability and return answer + retrieval data.
|
"""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.
|
are extracted from the state for downstream eval scoring.
|
||||||
|
|
||||||
The capability must produce a state with RAG-capability-shaped fields (citation
|
The capability must produce a state with RAG-capability-shaped fields (citation
|
||||||
index, searches, optional document filter) — i.e. ``RAGState`` or
|
index, searches, optional document filter and source scope) — i.e. ``RAGState``
|
||||||
``AnalysisState`` from ``haiku.rag.capabilities``.
|
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, deps, agent = _prepare_agent(
|
||||||
capability_factory,
|
capability_factory,
|
||||||
|
|
@ -187,6 +194,7 @@ async def run_capability_question(
|
||||||
capability_model,
|
capability_model,
|
||||||
document_filter,
|
document_filter,
|
||||||
request_limit,
|
request_limit,
|
||||||
|
sources=sources,
|
||||||
)
|
)
|
||||||
agent_result = await agent.run(question, deps=deps, message_history=message_history)
|
agent_result = await agent.run(question, deps=deps, message_history=message_history)
|
||||||
traffic = _count_tool_traffic(
|
traffic = _count_tool_traffic(
|
||||||
|
|
@ -205,6 +213,7 @@ async def run_capability_conversation(
|
||||||
capability_model: str | Model,
|
capability_model: str | Model,
|
||||||
document_filter: str | None = None,
|
document_filter: str | None = None,
|
||||||
compaction: bool = False,
|
compaction: bool = False,
|
||||||
|
sources: list[str] | None = None,
|
||||||
) -> list[CapabilityRunResult]:
|
) -> list[CapabilityRunResult]:
|
||||||
"""Run a conversation's user turns sequentially through one capability.
|
"""Run a conversation's user turns sequentially through one capability.
|
||||||
|
|
||||||
|
|
@ -223,6 +232,7 @@ async def run_capability_conversation(
|
||||||
document_filter=document_filter,
|
document_filter=document_filter,
|
||||||
request_limit=None,
|
request_limit=None,
|
||||||
compaction=compaction,
|
compaction=compaction,
|
||||||
|
sources=sources,
|
||||||
)
|
)
|
||||||
history: list[ModelMessage] | None = None
|
history: list[ModelMessage] | None = None
|
||||||
results: list[CapabilityRunResult] = []
|
results: list[CapabilityRunResult] = []
|
||||||
|
|
|
||||||
|
|
@ -345,6 +345,80 @@ async def test_conversation_applies_document_filter(tmp_path):
|
||||||
assert deps_seen[0].state["rag"]["document_filter"] == "uri = 'manual.pdf'"
|
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):
|
async def test_conversation_carries_one_state_dict_across_turns(tmp_path):
|
||||||
"""Capabilities read and write state through the deps dict; carrying the
|
"""Capabilities read and write state through the deps dict; carrying the
|
||||||
same dict across turns is what lets compaction see earlier questions'
|
same dict across turns is what lets compaction see earlier questions'
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue