Name the databases before the model runs

`ask(sources=["typo"])` reached the model, which discovered the name only if it
searched: requests spent on a selection that could never answer, and a run that
never searched answered anyway.

Checked by name, not by opening: a client covering a set opens a database when a
query reaches it, and validating by opening would open every one of them before
any search, letting a database nobody asked about fail the run.
This commit is contained in:
Yiorgis Gozadinos 2026-08-27 16:01:46 +03:00
parent 8791f12c37
commit 2fef67d9fd
No known key found for this signature in database
3 changed files with 67 additions and 0 deletions

View file

@ -817,6 +817,23 @@ class HaikuRAG:
found = await first_found(await self.clients_covering(), lookup)
return None if found is None else found[1]
def _require_known_sources(self, sources: "list[str] | None") -> None:
"""Fail on a name this client does not cover, opening nothing.
`clients_covering` answers the same question by opening the databases,
and a name is wrong whether or not what it names can be opened. `[]`
passes: a selection of nothing to search names nothing wrong.
"""
if sources is None:
return
covered = set(self.source_names)
unknown = [name for name in sources if name not in covered]
if unknown:
raise KeyError(
f"unknown database(s) {', '.join(sorted(set(unknown)))}; this "
f"client covers {', '.join(sorted(covered)) or 'a single unnamed database'}"
)
async def clients_covering(
self, sources: list[str] | None = None
) -> list["HaikuRAG"]:

View file

@ -63,6 +63,9 @@ async def ask(
)
from haiku.rag.utils import get_model
# Validate names without opening lazily covered databases.
client._require_known_sources(sources)
# No `db_path`: the lent client is what the capability reads through, and it
# already knows which databases that is.
capability = create_capability(
@ -124,6 +127,9 @@ async def analyze(
from haiku.rag.sandbox import AnalysisResult
from haiku.rag.utils import get_model
# Validate names without opening lazily covered databases.
client._require_known_sources(sources)
# No `db_path`: the lent client is what the capability reads through, and it
# already knows which databases that is.
capability = create_capability(

View file

@ -8,6 +8,7 @@ from haiku.rag.capabilities._tools import search_corpus
from haiku.rag.capabilities.rag import RAGState, create_capability
from haiku.rag.client import HaikuRAG
from haiku.rag.client.scope import DatabaseScope
from haiku.rag.client.session import FederatedSession
from haiku.rag.sandbox import AnalysisContext, Sandbox
from haiku.rag.store.models import SearchResult
from tests.multi_db.helpers import (
@ -200,6 +201,49 @@ class TestCollectionIdentityForTheModel:
assert result.stdout.count("['alpha', 'beta']") == 2
class TestNamingDatabasesBeforeTheModelRuns:
"""A name is checked at the boundary. Discovering it from a failed search
spends model requests, and a run can answer without reaching one."""
@pytest.mark.asyncio
async def test_ask_refuses_an_unknown_source_before_the_model(self, tmp_path):
config = _config(tmp_path, ["alpha", "beta"])
await _seed(config, "alpha", ["alpha document about cats"])
await _seed(config, "beta", ["beta document about cats"])
async with HaikuRAG(config=config) as rag:
with pytest.raises(KeyError, match="typo"):
await rag.ask("what about cats?", sources=["typo"])
@pytest.mark.asyncio
async def test_analyze_refuses_an_unknown_source_before_the_model(self, tmp_path):
config = _config(tmp_path, ["alpha", "beta"])
await _seed(config, "alpha", ["alpha document about cats"])
async with HaikuRAG(config=config) as rag:
with pytest.raises(KeyError, match="typo"):
await rag.analyze("how many?", sources=["typo"])
@pytest.mark.asyncio
async def test_checking_a_name_opens_nothing(self, tmp_path):
"""Opening to check would open every database on an unscoped question,
and let one nobody asked about fail a run before any search."""
config = _config(tmp_path, ["alpha", "beta"])
await _seed(config, "alpha", ["alpha document about cats"])
await _seed(config, "beta", ["beta document about cats"])
async with HaikuRAG(config=config) as rag:
assert isinstance(rag._session, FederatedSession)
rag._require_known_sources(None)
rag._require_known_sources(["alpha"])
rag._require_known_sources([])
with pytest.raises(KeyError, match="typo"):
rag._require_known_sources(["alpha", "typo"])
assert rag._session._sessions == {}
class TestLendingANamedClient:
@pytest.mark.asyncio
async def test_a_lent_named_client_names_the_citation(self, tmp_path):