diff --git a/tests/cassettes/test_multi_db_ask/TestAnalyzeAcrossDatabases.test_the_capability_searches_the_selected_databases.yaml b/tests/cassettes/multi_db/test_capabilities/TestAnalyzeAcrossDatabases.test_the_capability_searches_the_selected_databases.yaml similarity index 100% rename from tests/cassettes/test_multi_db_ask/TestAnalyzeAcrossDatabases.test_the_capability_searches_the_selected_databases.yaml rename to tests/cassettes/multi_db/test_capabilities/TestAnalyzeAcrossDatabases.test_the_capability_searches_the_selected_databases.yaml diff --git a/tests/cassettes/test_multi_db_ask/TestAskAcrossDatabases.test_searching_all_databases_reaches_both.yaml b/tests/cassettes/multi_db/test_capabilities/TestAskAcrossDatabases.test_searching_all_databases_reaches_both.yaml similarity index 100% rename from tests/cassettes/test_multi_db_ask/TestAskAcrossDatabases.test_searching_all_databases_reaches_both.yaml rename to tests/cassettes/multi_db/test_capabilities/TestAskAcrossDatabases.test_searching_all_databases_reaches_both.yaml diff --git a/tests/cassettes/test_multi_db_ask/TestAskAcrossDatabases.test_the_capability_searches_the_selected_databases.yaml b/tests/cassettes/multi_db/test_capabilities/TestAskAcrossDatabases.test_the_capability_searches_the_selected_databases.yaml similarity index 100% rename from tests/cassettes/test_multi_db_ask/TestAskAcrossDatabases.test_the_capability_searches_the_selected_databases.yaml rename to tests/cassettes/multi_db/test_capabilities/TestAskAcrossDatabases.test_the_capability_searches_the_selected_databases.yaml diff --git a/tests/cassettes/test_multi_db_ask/TestDatabaseIdentityForTheModel.test_in_code_search_names_the_database.yaml b/tests/cassettes/multi_db/test_capabilities/TestDatabaseIdentityForTheModel.test_in_code_search_names_the_database.yaml similarity index 100% rename from tests/cassettes/test_multi_db_ask/TestDatabaseIdentityForTheModel.test_in_code_search_names_the_database.yaml rename to tests/cassettes/multi_db/test_capabilities/TestDatabaseIdentityForTheModel.test_in_code_search_names_the_database.yaml diff --git a/tests/cassettes/test_multi_db_ask/TestStandaloneCapabilities.test_a_rag_capability_opens_the_configured_set.yaml b/tests/cassettes/multi_db/test_capabilities/TestStandaloneCapabilities.test_a_rag_capability_opens_the_configured_set.yaml similarity index 100% rename from tests/cassettes/test_multi_db_ask/TestStandaloneCapabilities.test_a_rag_capability_opens_the_configured_set.yaml rename to tests/cassettes/multi_db/test_capabilities/TestStandaloneCapabilities.test_a_rag_capability_opens_the_configured_set.yaml diff --git a/tests/cassettes/test_multi_db_ask/TestCiteFallback.test_an_id_from_a_selected_database_resolves_with_its_source.yaml b/tests/cassettes/multi_db/test_citations/TestCiteFallback.test_an_id_from_a_selected_database_resolves_with_its_source.yaml similarity index 100% rename from tests/cassettes/test_multi_db_ask/TestCiteFallback.test_an_id_from_a_selected_database_resolves_with_its_source.yaml rename to tests/cassettes/multi_db/test_citations/TestCiteFallback.test_an_id_from_a_selected_database_resolves_with_its_source.yaml diff --git a/tests/cassettes/test_multi_db_ask/TestCiteFallback.test_an_id_outside_the_selected_databases_does_not_resolve.yaml b/tests/cassettes/multi_db/test_citations/TestCiteFallback.test_an_id_outside_the_selected_databases_does_not_resolve.yaml similarity index 100% rename from tests/cassettes/test_multi_db_ask/TestCiteFallback.test_an_id_outside_the_selected_databases_does_not_resolve.yaml rename to tests/cassettes/multi_db/test_citations/TestCiteFallback.test_an_id_outside_the_selected_databases_does_not_resolve.yaml diff --git a/tests/multi_db/__init__.py b/tests/multi_db/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests/multi_db/conftest.py b/tests/multi_db/conftest.py new file mode 100644 index 00000000..e0593621 --- /dev/null +++ b/tests/multi_db/conftest.py @@ -0,0 +1,31 @@ +from pathlib import Path + +import pytest + +from haiku.rag.config import get_config + + +@pytest.fixture(scope="module") +def vcr_cassette_dir(request): + """Cassettes sit with the rest, under `tests/cassettes/multi_db/`.""" + module = request.module.__name__.rsplit(".", 1)[-1] + return str(Path(__file__).parent.parent / "cassettes" / "multi_db" / module) + + +@pytest.fixture +def query_embedding(monkeypatch): + """Vector search with no embedder behind it, recording the queries embedded. + + These tests are about which databases are asked and how often, not about + retrieval quality, and CI has no embedding endpoint. + """ + from haiku.rag.embeddings import EmbedderWrapper + + embedded: list[str] = [] + + async def embed_query(self, text): + embedded.append(text) + return [0.1] * get_config().embeddings.model.vector_dim + + monkeypatch.setattr(EmbedderWrapper, "embed_query", embed_query) + return embedded diff --git a/tests/multi_db/helpers.py b/tests/multi_db/helpers.py new file mode 100644 index 00000000..ca276fde --- /dev/null +++ b/tests/multi_db/helpers.py @@ -0,0 +1,96 @@ +"""Databases to run the multi-database tests against.""" + +from docling_core.types.doc.document import DoclingDocument +from docling_core.types.doc.labels import DocItemLabel + +from haiku.rag.client import HaikuRAG +from haiku.rag.config import get_config +from haiku.rag.config.models import AppConfig, LanceDBConfig +from haiku.rag.store.models import Chunk +from haiku.rag.utils import locate_database + + +def _config(tmp_path, names) -> AppConfig: + return AppConfig( + lancedb=LanceDBConfig( + databases={n: str(tmp_path / f"{n}.lancedb") for n in names} + ) + ) + + +async def _seed(config, name, contents): + """Precomputed embeddings and FTS queries keep the embedder out of the way: + these tests are about fusion, not retrieval quality.""" + dim = get_config().embeddings.model.vector_dim + async with HaikuRAG(config=config, create=True, sources=[name]) as rag: + for content in contents: + doc = DoclingDocument(name=content) + doc.add_text(label=DocItemLabel.TEXT, text=content) + await rag.import_document( + doc, + [Chunk(content=content, embedding=[0.1] * dim, order=0)], + uri=f"test://{name}/{content}", + ) + + +async def _restore_embedder(config, name, *, provider=None, model_name=None): + """Rewrite what one database records about the embedder that wrote it, + standing in for a database built elsewhere with another model.""" + import json + + import lancedb + + _, db_path = locate_database(config.lancedb.databases[name]) + assert db_path is not None + db = await lancedb.connect_async(str(db_path.resolve())) + table = await db.open_table("settings") + rows = ( + await table.query().where("id = 'settings'").limit(1).to_arrow() + ).to_pylist() + stored = json.loads(rows[0]["settings"]) + model = stored["embeddings"]["model"] + if provider is not None: + model["provider"] = provider + if model_name is not None: + model["name"] = model_name + await table.update({"settings": json.dumps(stored)}, where="id = 'settings'") + + +async def _seed_expandable(config, name, sentences): + """One document whose chunk covers a single item, so expansion has + neighbours to pull in and rebuilds the result rather than passing it + through.""" + dim = get_config().embeddings.model.vector_dim + doc = DoclingDocument(name=name) + for sentence in sentences: + doc.add_text(label=DocItemLabel.TEXT, text=sentence) + async with HaikuRAG(config=config, create=True, sources=[name]) as rag: + await rag.import_document( + doc, + [ + Chunk( + content=sentences[0], + embedding=[0.1] * dim, + order=0, + metadata={"doc_item_refs": ["#/texts/0"]}, + ) + ], + uri=f"test://{name}/expandable", + ) + + +class StubReranker: + """Scores the union, reversing it so the ordering is unmistakably its own.""" + + def __init__(self): + self.seen: list[str] = [] + + async def rerank(self, query, chunks, top_n): + self.seen = [c.content for c in chunks] + # Whatever the caller attached before handing them over. + self.attached = { + c.content.split()[0]: c._picture_data + for c in chunks + if getattr(c, "_picture_data", None) + } + return [(c, 1.0 - i) for i, c in enumerate(reversed(chunks))][:top_n] diff --git a/tests/multi_db/test_capabilities.py b/tests/multi_db/test_capabilities.py new file mode 100644 index 00000000..aefa2274 --- /dev/null +++ b/tests/multi_db/test_capabilities.py @@ -0,0 +1,207 @@ +"""Asking and analyzing across the databases a question covers.""" + +import pytest + +from haiku.rag.capabilities.rag import RAGState, create_capability +from haiku.rag.client import HaikuRAG +from haiku.rag.sandbox import AnalysisContext, Sandbox +from haiku.rag.store.models import SearchResult +from tests.multi_db.helpers import ( + _config, + _seed, +) + + +class TestAskAcrossDatabases: + @pytest.mark.asyncio + @pytest.mark.vcr() + async def test_the_capability_searches_the_selected_databases(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: + capability = create_capability(config=config, rag=rag, defer_loading=False) + capability.state = RAGState(sources=["alpha"]) + + formatted = await capability._search("cats", limit=10) + + assert isinstance(formatted, str) + assert "alpha" in formatted + assert "beta document" not in formatted + + @pytest.mark.asyncio + @pytest.mark.vcr() + async def test_searching_all_databases_reaches_both(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: + capability = create_capability(config=config, rag=rag, defer_loading=False) + capability.state = RAGState() + + formatted = await capability._search("cats", limit=10) + + assert isinstance(formatted, str) + assert "alpha document" in formatted + assert "beta document" in formatted + + +class TestStandaloneCapabilities: + """A capability nobody hands a client opens its own. It has to reach the + configured set, or a host that only registers capabilities gets one + database while the configuration names several.""" + + @pytest.mark.asyncio + @pytest.mark.vcr() + async def test_a_rag_capability_opens_the_configured_set(self, tmp_path): + from tests.capabilities.test_capabilities import Deps, make_context + + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha document about cats"]) + await _seed(config, "beta", ["beta document about cats"]) + + capability = create_capability(config=config, defer_loading=False) + assert capability.scope.names == ("alpha", "beta") + run = await capability.for_run(make_context(Deps())) + try: + formatted = await run._search("cats", limit=10) + finally: + await run._close() + + assert isinstance(formatted, str) + assert "alpha document" in formatted + assert "beta document" in formatted + + @pytest.mark.asyncio + @pytest.mark.vcr() + async def test_an_analysis_capability_mounts_the_configured_set(self, tmp_path): + from haiku.rag.capabilities.analysis import ( + create_capability as create_analysis, + ) + from tests.capabilities.test_capabilities import Deps, make_context + + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha document about cats"]) + await _seed(config, "beta", ["beta document about cats"]) + + capability = create_analysis(config=config, defer_loading=False) + run = await capability.for_run(make_context(Deps())) + try: + sandbox = await run._ensure_sandbox() + docs, owners = await sandbox._documents() + finally: + await run._close() + + assert len(docs) == 2 + assert {owner.source for owner in owners.values()} == {"alpha", "beta"} + + @pytest.mark.asyncio + async def test_a_single_configured_database_is_still_opened(self, tmp_path): + """One named database is a set of one, not a path to guess.""" + config = _config(tmp_path, ["alpha"]) + await _seed(config, "alpha", ["alpha document about cats"]) + + capability = create_capability(config=config, defer_loading=False) + rag = await capability._ensure_rag() + try: + assert rag.source == "alpha" + finally: + await capability._close() + + +class TestAnalyzeAcrossDatabases: + @pytest.mark.asyncio + @pytest.mark.vcr() + async def test_the_capability_searches_the_selected_databases(self, tmp_path): + """`analysis_search` is the same tool as the RAG one, and the sandbox is + scoped by the same selection.""" + from haiku.rag.capabilities.analysis import AnalysisState + from haiku.rag.capabilities.analysis import ( + create_capability as create_analysis, + ) + + 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: + capability = create_analysis(config=config, rag=rag, defer_loading=False) + capability.state = AnalysisState(sources=["alpha"]) + + formatted = await capability._search("cats", limit=10) + sandbox = await capability._ensure_sandbox() + await capability._close() + + assert isinstance(formatted, str) + assert "alpha document" in formatted + assert "beta document" not in formatted + assert sandbox._context.sources == ["alpha"] + + +class TestDatabaseIdentityForTheModel: + def test_a_result_names_its_database(self): + """The model has to attribute and compare evidence by database while it + composes the answer, not only afterwards through the citations.""" + result = SearchResult(content="body", score=0.9, source="alpha", chunk_id="c1") + + assert "Database: alpha" in result.format_for_agent() + + def test_an_unnamed_database_is_not_mentioned(self): + """A single unnamed database renders as it always has.""" + result = SearchResult(content="body", score=0.9, chunk_id="c1") + + assert "Database" not in result.format_for_agent() + + @pytest.mark.asyncio + @pytest.mark.vcr() + async def test_in_code_search_names_the_database(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: + sandbox = Sandbox( + db_path=None, + config=config, + context=AnalysisContext(), + rag=rag, + ) + try: + result = await sandbox.execute( + "rows = await search('cats', limit=10)\n" + "print(sorted(r['source'] for r in rows))\n" + "docs = await list_documents()\n" + "print(sorted(d['source'] for d in docs))" + ) + finally: + await sandbox.close() + + assert result.success, result.stderr + assert "['alpha', 'beta']" in result.stdout + assert result.stdout.count("['alpha', 'beta']") == 2 + + +class TestActionableFailures: + @pytest.mark.asyncio + async def test_a_migration_error_survives_being_named(self, tmp_path, temp_db_path): + """The remedy is the whole value of the message, and it names no location, + so it is not replaced by the database's name.""" + from haiku.rag.store.exceptions import MigrationRequiredError + + config = _config(tmp_path, ["alpha"]) + await _seed(config, "alpha", ["alpha document about cats"]) + + async with HaikuRAG(config=config, sources=["alpha"]) as rag: + await rag.store.set_haiku_version("0.20.0") + + with pytest.raises(MigrationRequiredError) as raised: + async with HaikuRAG(config=config, sources=["alpha"]): + pass + + # Both halves: which database failed, and what to run about it. + assert "haiku-rag migrate" in str(raised.value) + assert "alpha" in str(raised.value) + assert str(tmp_path) not in str(raised.value) diff --git a/tests/multi_db/test_citations.py b/tests/multi_db/test_citations.py new file mode 100644 index 00000000..19d7b21b --- /dev/null +++ b/tests/multi_db/test_citations.py @@ -0,0 +1,365 @@ +"""Citing evidence drawn from several databases.""" + +import pytest +from pydantic_ai import ModelRetry + +from haiku.rag.capabilities.rag import RAGState, create_capability +from haiku.rag.client import HaikuRAG +from haiku.rag.store.exceptions import ( + AmbiguousCitationError, +) +from haiku.rag.store.models import SearchResult +from haiku.rag.store.models.citation import Citation, resolve_citations +from tests.multi_db.helpers import ( + _config, + _seed, +) + + +class TestSharedChunkIds: + """A database copied from another holds the same chunk ids.""" + + @pytest.mark.asyncio + async def test_a_shared_id_does_not_confuse_the_fused_order(self, tmp_path): + """Arrival order breaks score ties, so it has to tell two databases' + identically-numbered chunks apart.""" + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one about cats"]) + await _seed(config, "beta", ["beta one about cats"]) + fused = [ + SearchResult(content="a0", score=0.5, chunk_id="a0", source="alpha"), + SearchResult(content="beta", score=0.5, chunk_id="shared", source="beta"), + SearchResult(content="alpha", score=0.5, chunk_id="shared", source="alpha"), + ] + + async with HaikuRAG(config=config) as rag: + expanded = await rag.expand_context(fused) + + assert [(r.source, r.chunk_id) for r in expanded] == [ + (r.source, r.chunk_id) for r in fused + ] + + def test_a_shared_id_cannot_be_cited(self): + """A citation records the id alone, so resolving one held by two + databases would attribute the answer to whichever came last.""" + results = [ + SearchResult( + content="alpha body", + score=0.9, + source="alpha", + chunk_id="c1", + document_id="d1", + document_uri="test://alpha/one", + ), + SearchResult( + content="beta body", + score=0.8, + source="beta", + chunk_id="c1", + document_id="d1", + document_uri="test://beta/one", + ), + ] + + with pytest.raises(AmbiguousCitationError, match="c1"): + resolve_citations(["c1"], results) + + def test_a_repeated_id_from_one_database_still_collapses(self): + """One database cannot hold two chunks under one id, so seeing it twice + is the same chunk seen twice.""" + results = [ + SearchResult( + content="first", + score=0.9, + source="alpha", + chunk_id="c1", + document_id="d1", + document_uri="test://alpha/one", + ), + SearchResult( + content="second", + score=0.8, + source="alpha", + chunk_id="c1", + document_id="d1", + document_uri="test://alpha/one", + ), + ] + + [citation] = resolve_citations(["c1"], results) + + assert citation.content == "first" + + def test_only_a_cited_id_has_to_be_unambiguous(self): + """An id the answer never cites attributes nothing.""" + shared = [ + SearchResult( + content=f"{name} body", + score=0.9, + source=name, + chunk_id="c1", + document_id="d1", + document_uri=f"test://{name}/one", + ) + for name in ("alpha", "beta") + ] + own = SearchResult( + content="alpha only", + score=0.7, + source="alpha", + chunk_id="c2", + document_id="d2", + document_uri="test://alpha/two", + ) + + [citation] = resolve_citations(["c2"], [*shared, own]) + + assert citation.source == "alpha" + + @pytest.mark.asyncio + async def test_an_unsearched_shared_id_is_refused_by_the_fallback(self, tmp_path): + """The direct lookup is the only place a collision shows for an id no + search returned, so it has to ask every database rather than take the + first that answers.""" + import shutil + + config = _config(tmp_path, ["alpha", "beta"]) + await _seed( + config, "alpha", ["alpha document about cats", "alpha on aardvarks"] + ) + shutil.copytree(tmp_path / "alpha.lancedb", tmp_path / "beta.lancedb") + + async with HaikuRAG(config=config) as rag: + alpha = (await rag.clients_for(["alpha"]))[0] + chunks = await alpha.chunk_repository.list_all() + [aardvark] = [c for c in chunks if "aardvark" in c.content] + assert aardvark.id is not None + + capability = create_capability(config=config, rag=rag, defer_loading=False) + capability.state = RAGState() + + # No search ran, so the id can only resolve through the fallback. + with pytest.raises(ModelRetry, match="more than one database"): + await capability._cite([aardvark.id]) + + @pytest.mark.asyncio + async def test_an_unsearched_id_in_one_database_still_resolves(self, tmp_path): + """The refusal is for a collision, not for looking through several + databases: an id only one of them holds still resolves.""" + from tests.capabilities.test_capabilities import Deps, make_context + + config = _config(tmp_path, ["alpha", "beta"]) + await _seed( + config, "alpha", ["alpha document about cats", "alpha on aardvarks"] + ) + await _seed(config, "beta", ["beta document about dogs"]) + + async with HaikuRAG(config=config) as rag: + alpha = (await rag.clients_for(["alpha"]))[0] + chunks = await alpha.chunk_repository.list_all() + [aardvark] = [c for c in chunks if "aardvark" in c.content] + assert aardvark.id is not None + + capability = create_capability(config=config, rag=rag, defer_loading=False) + run = await capability.for_run(make_context(Deps())) + + await run._cite([aardvark.id]) + + assert run.state is not None + [citation] = list(run.state.citation_index.values()) + assert citation.source == "alpha" + + @pytest.mark.asyncio + async def test_cite_asks_for_other_evidence(self, tmp_path): + capability = create_capability( + config=_config(tmp_path, ["alpha", "beta"]), defer_loading=False + ) + capability.state = RAGState( + searches={ + "cats": [ + SearchResult( + content=f"{name} body", + score=0.9, + source=name, + chunk_id="c1", + document_id="d1", + document_uri=f"test://{name}/one", + ) + for name in ("alpha", "beta") + ] + } + ) + + with pytest.raises(ModelRetry, match="appears once"): + await capability._cite(["c1"]) + + @pytest.mark.asyncio + async def test_cite_refuses_an_id_already_cited_from_another_database( + self, tmp_path + ): + """The citation index outlives the question, so the collision can arrive + a turn later than the search that would have shown it.""" + capability = create_capability( + config=_config(tmp_path, ["alpha", "beta"]), defer_loading=False + ) + capability.state = RAGState( + citation_index={ + "c1": Citation( + document_id="d1", + source="alpha", + chunk_id="c1", + document_uri="test://alpha/one", + content="alpha body", + ) + }, + searches={ + "cats": [ + SearchResult( + content="beta body", + score=0.9, + source="beta", + chunk_id="c1", + document_id="d1", + document_uri="test://beta/one", + ) + ] + }, + ) + + with pytest.raises(ModelRetry, match="another database"): + await capability._cite(["c1"]) + + +class TestCitationSource: + def test_a_citation_carries_the_result_source(self): + result = SearchResult( + content="body", + score=0.9, + source="alpha", + chunk_id="c1", + document_id="d1", + document_uri="test://alpha/one", + ) + + [citation] = resolve_citations(["c1"], [result]) + + assert citation.source == "alpha" + + def test_a_result_without_an_id_is_skipped(self): + """A result built by hand carries no chunk id, so nothing can cite it + and it takes part in no collision.""" + handmade = SearchResult(content="loose text", score=0.5) + real = SearchResult( + content="body", + score=0.9, + source="alpha", + chunk_id="c1", + document_id="d1", + document_uri="test://alpha/one", + ) + + [citation] = resolve_citations(["c1"], [handmade, real]) + + assert citation.chunk_id == "c1" + + def test_a_single_database_citation_has_no_source(self): + result = SearchResult( + content="body", + score=0.9, + chunk_id="c1", + document_id="d1", + document_uri="test://one", + ) + + [citation] = resolve_citations(["c1"], [result]) + + assert citation.source is None + + +class TestCiteFallback: + @pytest.mark.asyncio + @pytest.mark.vcr() + async def test_an_id_from_a_selected_database_resolves_with_its_source( + self, tmp_path + ): + """The fallback exists for a real id this run's searches did not return. + Across databases it looks through the selected ones and records which + held it.""" + from tests.capabilities.test_capabilities import Deps, make_context + + config = _config(tmp_path, ["alpha", "beta"]) + await _seed( + config, "alpha", ["alpha document about cats", "alpha on aardvarks"] + ) + await _seed(config, "beta", ["beta document about cats"]) + + async with HaikuRAG(config=config) as rag: + alpha = (await rag.clients_for(["alpha"]))[0] + chunks = await alpha.chunk_repository.list_all() + [aardvark] = [c for c in chunks if "aardvark" in c.content] + assert aardvark.id is not None + + capability = create_capability(config=config, rag=rag, defer_loading=False) + deps = Deps( + state={"rag": RAGState(sources=["alpha"]).model_dump(mode="json")} + ) + run = await capability.for_run(make_context(deps)) + # The search returns the cats chunk, never the aardvark one. + await run._search("cats", limit=10) + + await run._cite([aardvark.id]) + + assert run.state is not None + [citation] = list(run.state.citation_index.values()) + assert citation.chunk_id == aardvark.id + assert citation.source == "alpha" + + @pytest.mark.asyncio + @pytest.mark.vcr() + async def test_an_id_outside_the_selected_databases_does_not_resolve( + self, tmp_path + ): + """A question scoped to one database must not produce a citation from + another: the fallback looks only where the question looked.""" + from tests.capabilities.test_capabilities import Deps, make_context + + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha document about cats"]) + await _seed(config, "beta", ["beta document about dogs"]) + + async with HaikuRAG(config=config) as rag: + beta = (await rag.clients_for(["beta"]))[0] + [outside] = await beta.chunk_repository.list_all(limit=1) + assert outside.id is not None + + capability = create_capability(config=config, rag=rag, defer_loading=False) + deps = Deps( + state={"rag": RAGState(sources=["alpha"]).model_dump(mode="json")} + ) + run = await capability.for_run(make_context(deps)) + await run._search("cats", limit=10) + + with pytest.raises(ModelRetry): + await run._cite([outside.id]) + + @pytest.mark.asyncio + async def test_selecting_no_databases_cites_nothing(self, tmp_path): + """`sources=[]` selected nothing, which is not the same as everything: + the fallback must not go looking where the question never looked.""" + from tests.capabilities.test_capabilities import Deps, make_context + + 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: + alpha = (await rag.clients_for(["alpha"]))[0] + [chunk] = await alpha.chunk_repository.list_all(limit=1) + assert chunk.id is not None + + capability = create_capability(config=config, rag=rag, defer_loading=False) + deps = Deps(state={"rag": RAGState(sources=[]).model_dump(mode="json")}) + run = await capability.for_run(make_context(deps)) + + with pytest.raises(ModelRetry): + await run._cite([chunk.id]) diff --git a/tests/multi_db/test_documents.py b/tests/multi_db/test_documents.py new file mode 100644 index 00000000..9a7fe21d --- /dev/null +++ b/tests/multi_db/test_documents.py @@ -0,0 +1,232 @@ +"""Listing and looking up documents across databases.""" + +import pytest +from docling_core.types.doc.document import DoclingDocument +from docling_core.types.doc.labels import DocItemLabel + +from haiku.rag.client import HaikuRAG +from haiku.rag.config import get_config +from haiku.rag.store.models import Chunk +from tests.multi_db.helpers import ( + _config, + _seed, +) + + +class TestListingAcrossDatabases: + """The chat TUI's document filter lists documents through the client, and a + client covering a set has no repositories of its own.""" + + @pytest.mark.asyncio + async def test_listing_covers_every_database(self, tmp_path): + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one", "alpha two"]) + await _seed(config, "beta", ["beta one"]) + + async with HaikuRAG(config=config) as rag: + docs = await rag.list_documents() + + assert {d.uri for d in docs} == { + "test://alpha/alpha one", + "test://alpha/alpha two", + "test://beta/beta one", + } + + @pytest.mark.asyncio + async def test_counting_covers_every_database(self, tmp_path): + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one", "alpha two"]) + await _seed(config, "beta", ["beta one"]) + + async with HaikuRAG(config=config) as rag: + assert await rag.count_documents() == 3 + + @pytest.mark.asyncio + async def test_a_limit_bounds_the_merged_listing(self, tmp_path): + """A limit is that many documents in total, not that many per database.""" + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one", "alpha two"]) + await _seed(config, "beta", ["beta one", "beta two"]) + + async with HaikuRAG(config=config) as rag: + assert len(await rag.list_documents(limit=3)) == 3 + assert len(await rag.list_documents(limit=2, offset=2)) == 2 + assert len(await rag.list_documents(offset=3)) == 1 + + @pytest.mark.asyncio + async def test_a_page_shows_every_database(self, tmp_path): + """A window is taken across the databases, not filled from the first one: + concatenating hides every database after whichever was listed first.""" + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", [f"alpha {i}" for i in range(5)]) + await _seed(config, "beta", ["beta one"]) + + async with HaikuRAG(config=config) as rag: + page = await rag.list_documents(limit=3) + + assert len(page) == 3 + assert {(d.uri or "").split("/")[2] for d in page} == {"alpha", "beta"} + + @pytest.mark.asyncio + async def test_a_filter_reaches_every_database(self, tmp_path): + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one"]) + await _seed(config, "beta", ["beta one"]) + + async with HaikuRAG(config=config) as rag: + docs = await rag.list_documents(filter="uri LIKE 'test://beta/%'") + + assert [d.uri for d in docs] == ["test://beta/beta one"] + + +class TestLookupByIdentifier: + """An id or a URI says nothing about which database holds it, and a client + covering a set has no repositories of its own.""" + + @pytest.mark.asyncio + async def test_a_document_is_found_in_whichever_database_holds_it(self, tmp_path): + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one"]) + await _seed(config, "beta", ["beta one"]) + + async with HaikuRAG(config=config) as rag: + beta = (await rag.clients_for(["beta"]))[0] + [target] = await beta.document_repository.list_all(limit=1) + assert target.id is not None + + found = await rag.get_document_by_id(target.id) + by_uri = await rag.get_document_by_uri("test://alpha/alpha one") + resolved = await rag.resolve_document(target.id) + + assert found is not None and found.uri == "test://beta/beta one" + assert by_uri is not None and by_uri.uri == "test://alpha/alpha one" + assert resolved is not None and resolved.uri == "test://beta/beta one" + + @pytest.mark.asyncio + async def test_a_chunk_is_found_in_whichever_database_holds_it(self, tmp_path): + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one"]) + await _seed(config, "beta", ["beta one"]) + + async with HaikuRAG(config=config) as rag: + beta = (await rag.clients_for(["beta"]))[0] + [chunk] = await beta.chunk_repository.list_all(limit=1) + assert chunk.id is not None + + found = await rag.get_chunk_by_id(chunk.id) + + assert found is not None and found.content == "beta one" + + @pytest.mark.asyncio + async def test_a_document_held_by_two_databases_answers_from_the_first( + self, tmp_path + ): + """A database copied from another holds the same ids. A read has an + answer wherever it finds one, and which one it is has to be the + configured order rather than whichever replied first.""" + import shutil + + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one"]) + shutil.copytree(tmp_path / "alpha.lancedb", tmp_path / "beta.lancedb") + + async with HaikuRAG(config=config) as rag: + beta = (await rag.clients_for(["beta"]))[0] + [target] = await beta.document_repository.list_all(limit=1) + assert target.id is not None + + found = await rag.get_document_by_id(target.id) + + assert found is not None and found.source == "alpha" + + @pytest.mark.asyncio + async def test_an_unknown_identifier_is_absent_rather_than_an_error(self, tmp_path): + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one"]) + await _seed(config, "beta", ["beta one"]) + + async with HaikuRAG(config=config) as rag: + assert ( + await rag.get_document_by_id("00000000-0000-4000-8000-000000000000") + is None + ) + assert ( + await rag.get_chunk_by_id("00000000-0000-4000-8000-000000000000") + is None + ) + assert await rag.get_document_by_uri("test://nowhere") is None + + +class TestDocumentsNameTheirDatabase: + """A listing that spans databases is unreadable when the documents do not + say which one they came from, the same reason a search result carries one.""" + + @pytest.mark.asyncio + async def test_a_listing_names_each_document_s_database(self, tmp_path): + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one", "alpha two"]) + await _seed(config, "beta", ["beta one"]) + + async with HaikuRAG(config=config) as rag: + docs = await rag.list_documents() + + assert {d.uri: d.source for d in docs} == { + "test://alpha/alpha one": "alpha", + "test://alpha/alpha two": "alpha", + "test://beta/beta one": "beta", + } + + @pytest.mark.asyncio + async def test_a_looked_up_document_names_its_database(self, tmp_path): + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one"]) + await _seed(config, "beta", ["beta one"]) + + async with HaikuRAG(config=config) as rag: + beta = (await rag.clients_for(["beta"]))[0] + [target] = await beta.document_repository.list_all(limit=1) + assert target.id is not None + + by_id = await rag.get_document_by_id(target.id) + by_uri = await rag.get_document_by_uri("test://alpha/alpha one") + resolved = await rag.resolve_document(target.id) + + assert by_id is not None and by_id.source == "beta" + assert by_uri is not None and by_uri.source == "alpha" + assert resolved is not None and resolved.source == "beta" + + @pytest.mark.asyncio + async def test_one_named_database_still_names_itself(self, tmp_path): + """`haiku-rag --database alpha list` opens one database, and its name is + the whole reason the option exists.""" + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one"]) + await _seed(config, "beta", ["beta one"]) + + async with HaikuRAG(config=config, sources=["alpha"]) as rag: + [listed] = await rag.list_documents() + assert listed.id is not None + by_id = await rag.get_document_by_id(listed.id) + by_uri = await rag.get_document_by_uri("test://alpha/alpha one") + + assert listed.source == "alpha" + assert by_id is not None and by_id.source == "alpha" + assert by_uri is not None and by_uri.source == "alpha" + + @pytest.mark.asyncio + async def test_one_database_leaves_the_source_unset(self, tmp_path, temp_db_path): + """Nothing names the database when there is only one to name.""" + async with HaikuRAG(temp_db_path, create=True) as rag: + dim = get_config().embeddings.model.vector_dim + doc = DoclingDocument(name="solo") + doc.add_text(label=DocItemLabel.TEXT, text="solo") + await rag.import_document( + doc, + [Chunk(content="solo", embedding=[0.1] * dim, order=0)], + uri="test://solo", + ) + + [listed] = await rag.list_documents() + assert listed.source is None + assert listed.id is not None + assert (await rag.get_document_by_id(listed.id)).source is None diff --git a/tests/multi_db/test_expansion.py b/tests/multi_db/test_expansion.py new file mode 100644 index 00000000..023a24ae --- /dev/null +++ b/tests/multi_db/test_expansion.py @@ -0,0 +1,246 @@ +"""Expanding and enriching results through the database each came from.""" + +import pytest +from pydantic_ai import ModelRetry + +from haiku.rag.capabilities.rag import create_capability +from haiku.rag.client import HaikuRAG +from haiku.rag.store.models import Chunk, Document, DocumentItem, SearchResult +from tests.multi_db.helpers import ( + _config, + _seed, + _seed_expandable, +) + + +class TestExpansionRouting: + @pytest.mark.asyncio + async def test_expansion_routes_each_result_to_its_database(self, tmp_path): + """A federating client has no repositories of its own, so expansion has + to go through the database each result came from.""" + 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: + results = await rag.search("cats", search_type="fts", limit=10) + expanded = await rag.expand_context(results) + + assert {r.source for r in expanded} == {"alpha", "beta"} + for r in expanded: + assert r.source is not None + assert r.source in r.content + + @pytest.mark.asyncio + async def test_an_expanded_result_keeps_its_source(self, tmp_path): + """Expansion rebuilds the result, and the rebuilt one has to name the + database it was expanded through.""" + config = _config(tmp_path, ["alpha"]) + await _seed_expandable( + config, "alpha", ["cats sleep often", "cats also hunt", "cats purr"] + ) + + async with HaikuRAG(config=config) as rag: + results = await rag.search("cats", search_type="fts", limit=10) + expanded = await rag.expand_context(results) + + assert len(expanded) == 1 + assert "cats also hunt" in expanded[0].content, "expansion did not run" + assert expanded[0].source == "alpha" + + @pytest.mark.asyncio + async def test_a_federated_result_is_expanded_by_its_own_database(self, tmp_path): + """Routing is not enough: each result has to come back carrying the + neighbours of the database it was expanded through, and only those.""" + config = _config(tmp_path, ["alpha", "beta"]) + await _seed_expandable( + config, "alpha", ["cats sleep often", "alpha follows on"] + ) + await _seed_expandable(config, "beta", ["cats also hunt", "beta follows on"]) + + async with HaikuRAG(config=config) as rag: + results = await rag.search("cats", search_type="fts", limit=10) + expanded = await rag.expand_context(results) + + content = {r.source: r.content for r in expanded} + assert "alpha follows on" in content["alpha"] + assert "beta follows on" not in content["alpha"] + assert "beta follows on" in content["beta"] + + @pytest.mark.asyncio + async def test_expansion_keeps_tied_results_in_fused_order(self, tmp_path): + """Fused scores tie often, so grouping by database must not reorder + them: the tiebreak is the order they arrived in.""" + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one about cats", "alpha two about cats"]) + await _seed(config, "beta", ["beta one about cats"]) + + async with HaikuRAG(config=config) as rag: + found = await rag.search("cats", search_type="fts", limit=10) + by_source: dict[str, list[SearchResult]] = {} + for result in found: + by_source.setdefault(result.source or "", []).append(result) + # Interleaved, so grouping by database is visible as a reordering. + fused = [by_source["alpha"][0], by_source["beta"][0], by_source["alpha"][1]] + for result in fused: + result.score = 0.5 + + expanded = await rag.expand_context(fused) + + assert [r.chunk_id for r in expanded] == [r.chunk_id for r in fused] + + +class TestPictureDeduplication: + """One picture yields two chunks — a text-embedded one and an image-embedded + one — that collapse to the best. Two databases holding the same picture are + two results, not a duplicate.""" + + @staticmethod + def _picture(source, score): + return SearchResult( + content="a figure", + score=score, + source=source, + chunk_id=f"{source}-c", + document_id="doc-1", + doc_item_refs=["#/pictures/0"], + ) + + def test_the_same_picture_in_two_databases_survives(self): + from haiku.rag.client.search import _dedup_picture_chunks + + kept = _dedup_picture_chunks( + [self._picture("alpha", 0.9), self._picture("clone", 0.5)] + ) + + assert [r.source for r in kept] == ["alpha", "clone"] + + def test_duplicates_within_one_database_still_collapse(self): + from haiku.rag.client.search import _dedup_picture_chunks + + lower = self._picture("alpha", 0.5) + higher = self._picture("alpha", 0.9) + + kept = _dedup_picture_chunks([lower, higher]) + + assert kept == [higher] + + +class TestPictureRouting: + @pytest.mark.asyncio + async def test_a_picture_is_fetched_from_the_database_that_holds_it(self, tmp_path): + """A `self_ref` repeats across databases, so the citation's source is + what decides where the bytes come from.""" + 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: + beta = (await rag.clients_for(["beta"]))[0] + [document] = await beta.document_repository.list_all(limit=1) + assert document.id is not None + await beta.document_item_repository.create_all( + [ + DocumentItem( + document_id=document.id, + self_ref="#/pictures/0", + position=99, + label="picture", + text="", + picture_data=b"beta-picture", + ) + ] + ) + + assert ( + await rag.get_picture_bytes(document.id, "#/pictures/0", "beta") + == b"beta-picture" + ) + assert ( + await rag.get_picture_bytes(document.id, "#/pictures/0", "alpha") + is None + ) + + @pytest.mark.asyncio + async def test_a_single_database_needs_no_source(self, temp_db_path): + """One database is where the picture is, named or not.""" + async with HaikuRAG(temp_db_path, create=True) as rag: + document = await rag.document_repository.create( + Document(content="body", uri="test://one") + ) + assert document.id is not None + await rag.document_item_repository.create_all( + [ + DocumentItem( + document_id=document.id, + self_ref="#/pictures/0", + position=0, + label="picture", + text="", + picture_data=b"the-picture", + ) + ] + ) + + assert ( + await rag.get_picture_bytes(document.id, "#/pictures/0") + == b"the-picture" + ) + + @pytest.mark.asyncio + async def test_a_picture_lookup_without_a_source_is_refused(self, tmp_path): + """Federating, nothing can say which database holds an unqualified + reference.""" + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha document about cats"]) + + async with HaikuRAG(config=config) as rag: + with pytest.raises(ValueError, match="source"): + await rag.get_picture_bytes("doc-1", "#/pictures/0") + + +class TestFederatedEdges: + @pytest.mark.asyncio + async def test_expansion_passes_through_results_without_a_source(self, tmp_path): + """A caller can hand `expand_context` results it built itself. Those name + no database, so there is nowhere to expand them from.""" + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha document about cats"]) + await _seed(config, "beta", ["beta document about cats"]) + + handmade = SearchResult(content="handmade", score=0.4, doc_item_refs=[]) + + async with HaikuRAG(config=config) as rag: + found = await rag.search("cats", search_type="fts", limit=10) + expanded = await rag.expand_context([*found, handmade]) + + assert "handmade" in [r.content for r in expanded] + scores = [r.score for r in expanded] + assert scores == sorted(scores, reverse=True), "merged in score order" + + @pytest.mark.asyncio + async def test_a_chunk_without_a_document_is_not_cited(self, tmp_path): + """`Chunk.document_id` is optional, and a citation without a document has + nothing to point at.""" + from unittest.mock import AsyncMock, patch + + from haiku.rag.capabilities.rag import RAGCapability + from tests.capabilities.test_capabilities import ( + Deps, + _single_database_client, + make_context, + ) + + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha document about cats"]) + + orphan = _single_database_client() + orphan.get_chunk_by_id.return_value = Chunk( + id="orphan", document_id=None, content="no document" + ) + + capability = create_capability(config=config, defer_loading=False) + run = await capability.for_run(make_context(Deps())) + with patch.object(RAGCapability, "_ensure_rag", AsyncMock(return_value=orphan)): + with pytest.raises(ModelRetry): + await run._cite(["orphan"]) diff --git a/tests/multi_db/test_lifecycle.py b/tests/multi_db/test_lifecycle.py new file mode 100644 index 00000000..c179a837 --- /dev/null +++ b/tests/multi_db/test_lifecycle.py @@ -0,0 +1,475 @@ +"""Opening, borrowing and closing the databases a client covers.""" + +import asyncio + +import pytest +from docling_core.types.doc.document import DoclingDocument +from docling_core.types.doc.labels import DocItemLabel + +from haiku.rag.client import HaikuRAG +from haiku.rag.client.session import FederatedSession +from haiku.rag.config import get_config +from haiku.rag.store.exceptions import ( + AmbiguousDatabaseError, + SourceUnavailableError, +) +from haiku.rag.store.models import Chunk +from tests.multi_db.helpers import ( + _config, + _seed, +) + + +class TestOpeningDatabases: + @pytest.mark.asyncio + async def test_missing_databases_open_together(self, tmp_path): + """A cold fan-out costs one open, not their sum. On object storage a + serial loop is the difference between one round trip and N.""" + names = ["alpha", "beta", "gamma"] + config = _config(tmp_path, names) + for name in names: + await _seed(config, name, [f"{name} document about cats"]) + + async with HaikuRAG(config=config) as rag: + assert isinstance(rag._session, FederatedSession) + barrier = asyncio.Barrier(len(names)) + open_one = rag._session._open + + async def gated(ref): + # Every open has to be in flight before any of them finishes, so + # a serial loop cannot get past this and the wait times out. + await barrier.wait() + return await open_one(ref) + + rag._session._open = gated + clients = await asyncio.wait_for(rag.clients_for(names), timeout=15) + + assert {client.source for client in clients} == set(names) + + @pytest.mark.asyncio + async def test_a_failed_open_does_not_leak_the_ones_that_worked(self, tmp_path): + """Opening together means a failure has siblings already open. They are + tracked before it is reported, so closing the set closes them.""" + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha document about cats"]) + config.lancedb.databases["beta"] = str(tmp_path / "absent.lancedb") + + async with HaikuRAG(config=config) as rag: + with pytest.raises(SourceUnavailableError, match="beta"): + await rag.clients_for(["alpha", "beta"]) + + assert isinstance(rag._session, FederatedSession) + assert set(rag._session._sessions) == {"alpha"} + + @pytest.mark.asyncio + async def test_a_database_named_twice_is_opened_once(self, tmp_path): + """Fusion would count a repeated database as two rank lists.""" + 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: + clients = await rag.clients_for(["alpha", "alpha", "beta"]) + + assert [client.source for client in clients] == ["alpha", "beta"] + + @pytest.mark.asyncio + async def test_a_database_named_twice_returns_each_result_once(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: + results = await rag.search( + "cats", limit=10, search_type="fts", sources=["alpha", "alpha"] + ) + + assert [r.source for r in results] == ["alpha"] + + @pytest.mark.asyncio + async def test_one_database_named_twice_is_still_that_database(self, tmp_path): + """A client covering a single named database compares the selection + against its own name, so repeats have to collapse first.""" + config = _config(tmp_path, ["alpha"]) + await _seed(config, "alpha", ["alpha document about cats"]) + + async with HaikuRAG(config=config) as rag: + covering = await rag.clients_covering(["alpha", "alpha"]) + + assert [client.source for client in covering] == ["alpha"] + + +class TestClosingASet: + @pytest.mark.asyncio + async def test_every_database_opened_is_released(self, tmp_path): + """A covered database owns an embedder and may owe a vacuum. Closing only + its connection would leave both behind.""" + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one"]) + await _seed(config, "beta", ["beta one"]) + + released: list[str | None] = [] + drained: list[str | None] = [] + + async with HaikuRAG(config=config, read_only=True) as rag: + assert isinstance(rag._session, FederatedSession) + await rag.clients_for(["alpha", "beta"]) + for name, session in rag._session._sessions.items(): + original = session.store.embedder.aclose + drain = session.drain_vacuum + + async def release(_original=original, _name=name): + released.append(_name) + return await _original() + + async def drain_it(_drain=drain, _name=name): + drained.append(_name) + return await _drain() + + session.store.embedder.aclose = release + session.drain_vacuum = drain_it + + assert sorted(released) == ["alpha", "beta"] + assert sorted(drained) == ["alpha", "beta"] + + +class TestBorrowedDatabases: + """A client for one of a set wraps a database the set opened.""" + + @pytest.mark.asyncio + async def test_closing_a_borrowed_client_leaves_the_set_working(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: + (alpha,) = await rag.clients_for(["alpha"]) + store = alpha.store + + alpha.close() + assert store.db.is_open(), "close() closed a database it borrowed" + + await alpha.__aexit__(None, None, None) + assert store.db.is_open(), "exit closed a database it borrowed" + + results = await rag.search("cats", search_type="fts") + + assert {r.source for r in results} == {"alpha", "beta"} + assert not store.db.is_open(), "the set left a database open" + + @pytest.mark.asyncio + async def test_entering_a_borrowed_client_reuses_its_database(self, tmp_path): + """`async with` on a borrowed client is a plausible thing to write. + Opening a second session would leak it, since teardown declines to close + what this client did not open.""" + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha document about cats"]) + + async with HaikuRAG(config=config) as rag: + (alpha,) = await rag.clients_for(["alpha"]) + borrowed = alpha.store + + async with alpha as entered: + assert entered is alpha + assert alpha.store is borrowed, "entry opened a second database" + + assert borrowed.db.is_open(), "exit closed a database it borrowed" + assert alpha.store is borrowed + + assert not borrowed.db.is_open(), "the set left a database open" + + @pytest.mark.asyncio + async def test_a_borrowed_client_releases_what_it_built(self, tmp_path): + """Its reranker is its own; the database it wraps is not.""" + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha document about cats"]) + + closed: list[str] = [] + + class Reranker: + async def aclose(self): + closed.append("reranker") + + async with HaikuRAG(config=config) as rag: + (alpha,) = await rag.clients_for(["alpha"]) + alpha.__dict__["reranker"] = Reranker() + + assert closed == ["reranker"] + + +class TestLazyOpening: + @pytest.mark.asyncio + async def test_entering_opens_nothing(self, tmp_path): + """25 configured databases queried a few at a time must not all open.""" + 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 rag._clients == {} + + @pytest.mark.asyncio + async def test_only_the_selected_database_opens(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: + await rag.search("cats", search_type="fts", sources=["alpha"]) + assert list(rag._clients) == ["alpha"] + + @pytest.mark.asyncio + async def test_an_unselected_broken_database_does_not_break_a_query(self, tmp_path): + """A database nobody asked for cannot fail a query.""" + config = _config(tmp_path, ["alpha", "missing"]) + await _seed(config, "alpha", ["alpha document about cats"]) + + async with HaikuRAG(config=config) as rag: + results = await rag.search("cats", search_type="fts", sources=["alpha"]) + + assert [r.source for r in results] == ["alpha"] + + @pytest.mark.asyncio + async def test_a_selected_broken_database_fails_the_query(self, tmp_path): + config = _config(tmp_path, ["alpha", "missing"]) + await _seed(config, "alpha", ["alpha document about cats"]) + + async with HaikuRAG(config=config) as rag: + with pytest.raises(SourceUnavailableError, match="missing"): + await rag.search("cats", search_type="fts") + + +class TestReadOnlyMode: + @pytest.mark.asyncio + async def test_a_client_covering_a_set_reports_its_mode(self, tmp_path): + """A client covering a set has no store of its own to ask.""" + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one"]) + await _seed(config, "beta", ["beta one"]) + + async with HaikuRAG(config=config, read_only=True) as rag: + assert rag.is_read_only is True + async with HaikuRAG(config=config) as rag: + assert rag.is_read_only is False + + +class TestFailureNaming: + @pytest.mark.asyncio + async def test_a_single_named_database_is_reported_by_name(self, tmp_path): + """One configured database is still a named one: it must not fall back to + the raw error, which spells out the path.""" + config = _config(tmp_path, ["alpha"]) + + with pytest.raises(SourceUnavailableError, match="alpha") as caught: + async with HaikuRAG(config=config): + pass + + assert str(tmp_path) not in str(caught.value) + assert caught.value.__cause__ is None + + @pytest.mark.asyncio + async def test_a_legacy_uri_client_keeps_its_error(self, tmp_path): + """Nothing named it, so there is no name to report instead.""" + with pytest.raises(FileNotFoundError): + async with HaikuRAG(tmp_path / "nope.lancedb"): + pass + + @pytest.mark.asyncio + async def test_the_location_is_absent_from_the_whole_chain(self, tmp_path): + config = _config(tmp_path, ["alpha", "missing"]) + await _seed(config, "alpha", ["alpha document about cats"]) + + with pytest.raises(SourceUnavailableError) as caught: + async with HaikuRAG(config=config) as rag: + await rag.search("cats", search_type="fts") + + rendered = str(caught.value) + error = caught.value.__cause__ or caught.value.__context__ + assert "missing.lancedb" not in rendered + assert error is None, "the location-bearing cause is still attached" + + +class TestCreatingNeedsOneDatabase: + """Creating names a database. Covering a set, the flag had nothing to act on + and was accepted anyway, leaving the first query to fail on whichever + database turned out to be missing.""" + + @pytest.mark.asyncio + async def test_creating_a_set_is_refused(self, tmp_path): + config = _config(tmp_path, ["alpha", "beta"]) + + with pytest.raises(AmbiguousDatabaseError, match="alpha, beta"): + async with HaikuRAG(config=config, create=True): + pass + + @pytest.mark.asyncio + async def test_naming_one_of_the_set_creates_it(self, tmp_path): + config = _config(tmp_path, ["alpha", "beta"]) + + async with HaikuRAG(config=config, create=True, sources=["alpha"]) as rag: + assert await rag.count_documents() == 0 + + assert (tmp_path / "alpha.lancedb").exists() + assert not (tmp_path / "beta.lancedb").exists() + + @pytest.mark.asyncio + async def test_covering_a_set_without_creating_is_unaffected(self, tmp_path): + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one"]) + await _seed(config, "beta", ["beta one"]) + + async with HaikuRAG(config=config) as rag: + assert await rag.count_documents() == 2 + + +class TestOperationsThatNeedOneDatabase: + @pytest.mark.asyncio + async def test_writing_names_the_databases_it_covers(self, tmp_path): + """A domain error, so a caller can tell an unsupported selection from a + missing attribute.""" + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one"]) + await _seed(config, "beta", ["beta one"]) + + async with HaikuRAG(config=config) as rag: + with pytest.raises(AmbiguousDatabaseError, match="alpha, beta"): + await rag.create_document("orphan") + with pytest.raises(AmbiguousDatabaseError, match="clients_for"): + await rag.vacuum() + with pytest.raises(AmbiguousDatabaseError, match="close"): + rag.close() + + @pytest.mark.asyncio + async def test_a_set_has_no_store_of_its_own(self, tmp_path): + """A store and its repositories belong to one database. `clients_for` + reaches the one holding a given database.""" + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one"]) + + async with HaikuRAG(config=config) as rag: + for name in ( + "store", + "document_repository", + "chunk_repository", + "document_item_repository", + ): + with pytest.raises(AttributeError, match=name): + getattr(rag, name) + + @pytest.mark.asyncio + async def test_a_selected_database_is_still_writable(self, tmp_path): + """Naming one of the set is how a write picks its database.""" + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one"]) + await _seed(config, "beta", ["beta one"]) + + dim = get_config().embeddings.model.vector_dim + written = DoclingDocument(name="written") + written.add_text(label=DocItemLabel.TEXT, text="written") + + async with HaikuRAG(config=config) as rag: + alpha = (await rag.clients_for(["alpha"]))[0] + assert alpha.is_read_only is False + document = await alpha.import_document( + written, + [Chunk(content="written", embedding=[0.1] * dim, order=0)], + uri="test://alpha/written", + ) + assert await alpha.count_documents() == 2 + + assert document.id is not None + + +class TestDatabaseIndependentWork: + """Converting, chunking and titling are functions of the configuration, not + of a database, so covering a set does not stop them.""" + + @pytest.mark.asyncio + async def test_chunking_opens_no_database(self, tmp_path, monkeypatch): + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one"]) + await _seed(config, "beta", ["beta one"]) + + opened: list[str] = [] + + async def refuse(self, ref): + opened.append(ref.name) + raise AssertionError("opened a database to chunk a document") + + monkeypatch.setattr(FederatedSession, "_open", refuse) + + doc = DoclingDocument(name="note") + doc.add_text( + label=DocItemLabel.TEXT, text="Boltzmann machines are energy based." + ) + + async with HaikuRAG(config=config, read_only=True) as rag: + chunks = await rag.chunk(doc) + + assert opened == [] + assert [c.content for c in chunks] + + @pytest.mark.asyncio + async def test_the_embedder_is_built_once_and_closed_once( + self, tmp_path, monkeypatch + ): + """The parent owns the embedder it built, so leaving the context closes + it, once.""" + from haiku.rag.embeddings import EmbedderWrapper + + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one"]) + + closed: list[object] = [] + original = EmbedderWrapper.aclose + + async def counting(self): + closed.append(self) + return await original(self) + + monkeypatch.setattr(EmbedderWrapper, "aclose", counting) + + rag = HaikuRAG(config=config, read_only=True) + async with rag: + built = rag.embedder + assert rag.embedder is built + + assert closed == [built] + + @pytest.mark.asyncio + async def test_re_entering_a_set_builds_a_fresh_embedder(self, tmp_path): + """Teardown closes the embedder, so keeping it would hand the next + context one that is already closed.""" + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one"]) + + rag = HaikuRAG(config=config, read_only=True) + async with rag: + first = rag.embedder + async with rag: + assert rag.embedder is not first + + @pytest.mark.asyncio + async def test_re_entering_one_database_builds_a_fresh_embedder(self, temp_db_path): + """One database opens a new store on re-entry, and the embedder is that + store's.""" + rag = HaikuRAG(temp_db_path, create=True) + async with rag: + first = rag.embedder + async with rag: + assert rag.embedder is rag.store.embedder + assert rag.embedder is not first + + @pytest.mark.asyncio + async def test_a_set_nobody_asked_anything_of_builds_no_embedder(self, tmp_path): + """Built on first use, so a client that answered nothing holds nothing.""" + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one"]) + + async with HaikuRAG(config=config, read_only=True) as rag: + assert "embedder" not in rag.__dict__ + + @pytest.mark.asyncio + async def test_one_database_still_uses_its_store_s_embedder(self, temp_db_path): + async with HaikuRAG(temp_db_path, create=True) as rag: + assert rag.embedder is rag.store.embedder diff --git a/tests/multi_db/test_scope.py b/tests/multi_db/test_scope.py new file mode 100644 index 00000000..a081b18f --- /dev/null +++ b/tests/multi_db/test_scope.py @@ -0,0 +1,288 @@ +"""Resolving which databases an operation covers.""" + +import pytest +from pydantic import ValidationError + +from haiku.rag.client import HaikuRAG +from haiku.rag.client.scope import DatabaseScope +from haiku.rag.config.models import AppConfig, LanceDBConfig +from haiku.rag.utils import locate_database +from tests.multi_db.helpers import ( + _config, + _seed, +) + + +class TestConfig: + def test_databases_and_uri_are_mutually_exclusive(self): + with pytest.raises(ValidationError, match="databases"): + LanceDBConfig( + uri="s3://b/one.lancedb", databases={"one": "s3://b/one.lancedb"} + ) + + def test_databases_alone_is_fine(self): + config = LanceDBConfig(databases={"one": "s3://b/one.lancedb"}) + assert config.databases == {"one": "s3://b/one.lancedb"} + + def test_uri_alone_is_fine(self): + assert LanceDBConfig(uri="s3://b/one.lancedb").databases == {} + + +class TestNamingIsRequired: + def test_a_blank_name_is_rejected(self): + """An unnamed database is unreachable: every source check reads the + empty name as no name at all.""" + with pytest.raises(ValidationError, match="entry with no name"): + LanceDBConfig(databases={"": "/tmp/a.lancedb"}) + with pytest.raises(ValidationError, match="entry with no name"): + LanceDBConfig(databases={" ": "/tmp/a.lancedb"}) + + def test_a_blank_location_is_rejected(self): + """A blank location resolves to the working directory.""" + with pytest.raises( + ValidationError, match=r"databases\[alpha\] has no location" + ): + LanceDBConfig(databases={"alpha": ""}) + + +class TestNamingADatabaseDirectly: + @pytest.mark.asyncio + async def test_an_explicit_db_path_wins_over_the_configured_set( + self, tmp_path, temp_db_path + ): + """A caller that names a path means that database, not the configured + set: the CLI resolves `--db` to one and must not fan out instead.""" + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha document about cats"]) + + async with HaikuRAG(temp_db_path, config=config, create=True) as rag: + assert not rag.covers_multiple + assert rag.source is None + assert rag.store.db_path == temp_db_path + + @pytest.mark.asyncio + async def test_one_configured_database_is_opened_by_name(self, tmp_path): + """A set of one is not federated, and the client resolves it.""" + config = _config(tmp_path, ["alpha"]) + await _seed(config, "alpha", ["alpha document about cats"]) + + async with HaikuRAG(config=config) as rag: + assert not rag.covers_multiple + assert rag.source == "alpha" + results = await rag.search("cats", search_type="fts", limit=10) + + assert [r.source for r in results] == ["alpha"] + + +class TestOneConfiguredLocation: + """`lancedb.uri` places one unnamed database, at a URI or at a local path.""" + + def _config(self, location) -> AppConfig: + return AppConfig(lancedb=LanceDBConfig(uri=str(location))) + + @pytest.mark.asyncio + async def test_a_local_uri_opens_the_configured_database(self, tmp_path): + located = tmp_path / "notes.lancedb" + config = self._config(located) + + async with HaikuRAG(config=config, create=True) as rag: + assert rag.store.db_path == located + # It places a database without naming one: only `lancedb.databases` + # assigns the name results and citations carry. + assert rag.source is None + assert located.exists() + + @pytest.mark.asyncio + async def test_an_explicit_path_overrides_a_local_uri(self, tmp_path): + """`--db` overrides the configured location for one invocation.""" + config = self._config(tmp_path / "configured.lancedb") + chosen = tmp_path / "chosen.lancedb" + + async with HaikuRAG(chosen, config=config, create=True) as rag: + assert rag.store.db_path == chosen + assert chosen.exists() + assert not (tmp_path / "configured.lancedb").exists() + + @pytest.mark.asyncio + async def test_a_local_uri_that_does_not_exist_is_refused(self, tmp_path): + """A mistyped path fails instead of quietly becoming an empty database, + which is what a value carrying a scheme would do.""" + config = self._config(tmp_path / "typo.lancedb") + + with pytest.raises(FileNotFoundError): + async with HaikuRAG(config=config): + pass + assert not (tmp_path / "typo.lancedb").exists() + + def test_a_uri_with_a_scheme_stays_a_uri(self, tmp_path): + """Object storage has no local path to check, and a location that does + not exist yet is normal there.""" + from haiku.rag.store.engine import ConnectionMode + + config = self._config("s3://bucket/one.lancedb") + + [ref] = DatabaseScope.resolve(config).databases + one, db_path = ref.connection(config) + + assert db_path is None + assert ConnectionMode.from_config(one) == ConnectionMode.OBJECT_STORAGE + + +class TestLocate: + def test_a_scheme_is_a_uri(self): + assert locate_database("s3://bucket/one.lancedb") == ( + "s3://bucket/one.lancedb", + None, + ) + + def test_anything_else_is_a_local_path(self): + uri, db_path = locate_database("/data/one.lancedb") + assert uri == "" + assert db_path is not None and str(db_path) == "/data/one.lancedb" + + +class TestSelection: + @pytest.mark.asyncio + async def test_unknown_source_at_construction_is_rejected(self, tmp_path): + config = _config(tmp_path, ["alpha", "beta"]) + + with pytest.raises(KeyError, match="nope"): + async with HaikuRAG(config=config, sources=["nope"]): + pass + + @pytest.mark.asyncio + async def test_unknown_source_across_several_databases_is_rejected(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="nope"): + await rag.search("cats", search_type="fts", sources=["nope"]) + + @pytest.mark.asyncio + async def test_no_matches_anywhere_returns_nothing(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: + assert await rag.search("aardvarks", search_type="fts") == [] + + +class TestPlacingADatabase: + """What a client says about the databases it covers, so nothing outside has + to read its private state to find out.""" + + @pytest.mark.asyncio + async def test_a_set_names_every_database_it_covers(self, tmp_path): + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one"]) + await _seed(config, "beta", ["beta one"]) + + async with HaikuRAG(config=config, read_only=True) as rag: + assert rag.covers_multiple + assert rag.source_names == ("alpha", "beta") + assert rag.source is None + + @pytest.mark.asyncio + async def test_one_named_database_names_itself(self, tmp_path): + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one"]) + + async with HaikuRAG(config=config, read_only=True, sources=["alpha"]) as rag: + assert not rag.covers_multiple + assert rag.source_names == ("alpha",) + assert rag.source == "alpha" + + @pytest.mark.asyncio + async def test_a_named_database_keeps_its_name_on_re_entry(self, tmp_path): + """Entering derives a single-database configuration from what was + configured. Deriving it from the last derivation loses the name.""" + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha document about cats"]) + + rag = HaikuRAG(config=config, read_only=True, sources=["alpha"]) + async with rag: + assert rag.source == "alpha" + async with rag: + assert rag.source == "alpha" + assert rag.source_names == ("alpha",) + results = await rag.search("cats", search_type="fts") + + assert {r.source for r in results} == {"alpha"} + + @pytest.mark.asyncio + async def test_an_unnamed_database_names_nothing(self, temp_db_path): + async with HaikuRAG(temp_db_path, create=True) as rag: + assert rag.source_names == () + assert rag.source is None + + @pytest.mark.asyncio + async def test_the_reader_for_a_database_is_the_client_holding_it(self, tmp_path): + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one"]) + await _seed(config, "beta", ["beta one"]) + + async with HaikuRAG(config=config, read_only=True) as rag: + reader = await rag.reader_for("beta") + + assert reader is not None + assert reader.source == "beta" + # Asked twice, the same wrapper comes back. + assert await rag.reader_for("beta") is reader + + @pytest.mark.asyncio + async def test_a_client_reading_one_database_is_its_own_reader(self, temp_db_path): + async with HaikuRAG(temp_db_path, create=True) as rag: + assert await rag.reader_for(None) is rag + assert await rag.reader_for("anything") is rag + + @pytest.mark.asyncio + async def test_a_set_cannot_place_evidence_that_names_no_database(self, tmp_path): + """Evidence recorded before databases could be named carries no source.""" + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one"]) + + async with HaikuRAG(config=config, read_only=True) as rag: + assert await rag.reader_for(None) is None + + +class TestNamingOneOfTheSetOnTheCommandLine: + """`--database NAME` reaches the application layer as a name, and every + client it opens has to honour it — one that ignores it covers the set and + quietly answers from the wrong database.""" + + @pytest.mark.asyncio + async def test_a_named_database_is_the_one_read(self, tmp_path, capsys): + from haiku.rag.app import HaikuRAGApp + + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha document about cats"]) + await _seed(config, "beta", ["beta document about cats"]) + + scope = DatabaseScope.resolve(config).select(["beta"]) + app = HaikuRAGApp(scope=scope, config=config, read_only=True) + await app.list_documents() + + # Rich wraps long lines, so match the unwrapped part of the URI. + printed = capsys.readouterr().out + assert "test://beta/" in printed + assert "test://alpha/" not in printed + + @pytest.mark.asyncio + async def test_naming_none_of_them_covers_the_set(self, tmp_path, capsys): + from haiku.rag.app import HaikuRAGApp + + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha document about cats"]) + await _seed(config, "beta", ["beta document about cats"]) + + app = HaikuRAGApp( + scope=DatabaseScope.resolve(config), config=config, read_only=True + ) + await app.list_documents() + + printed = capsys.readouterr().out + assert "test://alpha/" in printed + assert "test://beta/" in printed diff --git a/tests/multi_db/test_search.py b/tests/multi_db/test_search.py new file mode 100644 index 00000000..bb7984dc --- /dev/null +++ b/tests/multi_db/test_search.py @@ -0,0 +1,375 @@ +"""Searching several databases and fusing what they return.""" + +import pytest +from docling_core.types.doc.document import DoclingDocument +from docling_core.types.doc.labels import DocItemLabel + +from haiku.rag.client import HaikuRAG +from haiku.rag.client.session import FederatedSession +from haiku.rag.config import get_config +from haiku.rag.store.exceptions import ( + ConfigMismatchError, + SourceUnavailableError, +) +from haiku.rag.store.models import Chunk, DocumentItem +from tests.multi_db.helpers import ( + StubReranker, + _config, + _restore_embedder, + _seed, +) + + +class TestFederatedSearch: + @pytest.mark.asyncio + async def test_results_carry_their_source(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: + results = await rag.search("cats", limit=10, search_type="fts") + + assert {r.source for r in results} == {"alpha", "beta"} + for r in results: + assert r.source is not None + assert r.source in r.content + + @pytest.mark.asyncio + async def test_sources_selects_a_subset(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: + results = await rag.search( + "cats", limit=10, search_type="fts", sources=["alpha"] + ) + + assert {r.source for r in results} == {"alpha"} + + @pytest.mark.asyncio + async def test_unknown_source_is_rejected(self, tmp_path): + config = _config(tmp_path, ["alpha"]) + await _seed(config, "alpha", ["alpha document about cats"]) + + async with HaikuRAG(config=config) as rag: + with pytest.raises(KeyError, match="nope"): + await rag.search("cats", search_type="fts", sources=["nope"]) + + @pytest.mark.asyncio + async def test_an_unopenable_database_fails_the_query(self, tmp_path): + config = _config(tmp_path, ["alpha", "missing"]) + await _seed(config, "alpha", ["alpha document about cats"]) + + with pytest.raises(SourceUnavailableError, match="missing"): + async with HaikuRAG(config=config) as rag: + await rag.search("cats", search_type="fts") + + +class TestSingleDatabaseUnchanged: + @pytest.mark.asyncio + async def test_source_is_unset_without_configured_databases(self, temp_db_path): + async with HaikuRAG(temp_db_path, create=True) as rag: + doc = DoclingDocument(name="one") + doc.add_text(label=DocItemLabel.TEXT, text="a document about cats") + await rag.import_document( + doc, + [ + Chunk( + content="a document about cats", + embedding=[0.1] * get_config().embeddings.model.vector_dim, + order=0, + ) + ], + uri="test://one", + ) + results = await rag.search("cats", search_type="fts") + + assert results + assert all(r.source is None for r in results) + + +class TestOneQueryVector: + @pytest.mark.asyncio + async def test_a_search_embeds_the_query_once_for_the_whole_set( + self, tmp_path, query_embedding + ): + """Each database owns an embedder, so embedding per database costs a + round trip each on a remote endpoint.""" + config = _config(tmp_path, ["alpha", "beta", "gamma"]) + for name in ("alpha", "beta", "gamma"): + await _seed(config, name, [f"{name} one"]) + + async with HaikuRAG(config=config, read_only=True) as rag: + await rag.search("one") + + assert query_embedding == ["one"] + + +class TestOneEmbedderAcrossTheSet: + """A set is searched with one query vector, so a database written with + another model would answer from a different space.""" + + @pytest.mark.asyncio + async def test_disagreeing_databases_cannot_be_searched_together(self, tmp_path): + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one"]) + await _seed(config, "beta", ["beta one"]) + await _restore_embedder(config, "beta", model_name="some-other-model") + + async with HaikuRAG(config=config, read_only=True) as rag: + with pytest.raises(ConfigMismatchError, match="different embedders"): + await rag.search("one") + + @pytest.mark.asyncio + async def test_a_database_asked_for_alone_is_never_compared( + self, tmp_path, query_embedding + ): + """Only databases searched together have to agree.""" + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one"]) + await _seed(config, "beta", ["beta one"]) + await _restore_embedder(config, "beta", model_name="some-other-model") + + async with HaikuRAG(config=config, read_only=True) as rag: + assert await rag.search("one", sources=["alpha"]) is not None + assert await rag.count_documents(filter=None) is not None + + @pytest.mark.asyncio + async def test_full_text_search_needs_no_agreement(self, tmp_path): + """Full-text search embeds nothing, so which model wrote each database + does not come into it.""" + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one"]) + await _seed(config, "beta", ["beta one"]) + await _restore_embedder(config, "beta", model_name="some-other-model") + + async with HaikuRAG(config=config, read_only=True) as rag: + results = await rag.search("one", search_type="fts") + + assert {r.source for r in results} == {"alpha", "beta"} + + @pytest.mark.asyncio + async def test_agreeing_databases_search_together(self, tmp_path, query_embedding): + """The databases agree with each other; that they were written by a + differently-spelled provider than the config is the soft case.""" + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha one"]) + await _seed(config, "beta", ["beta one"]) + await _restore_embedder(config, "alpha", provider="openai") + await _restore_embedder(config, "beta", provider="openai") + + async with HaikuRAG(config=config, read_only=True) as rag: + assert len(await rag.search("one")) > 0 + + +class TestRerankerFusion: + @pytest.mark.asyncio + async def test_the_reranker_scores_the_union_and_owners_survive( + self, tmp_path, monkeypatch + ): + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha document about cats"]) + await _seed(config, "beta", ["beta document about cats"]) + + stub = StubReranker() + monkeypatch.setattr(HaikuRAG, "reranker", property(lambda self: stub)) + + async with HaikuRAG(config=config) as rag: + results = await rag.search("cats", limit=2, search_type="fts") + + # It saw both databases' candidates, not one database at a time. + assert len(stub.seen) == 2 + assert {c.split()[0] for c in stub.seen} == {"alpha", "beta"} + # Each result still knows which database it came from. + for r in results: + assert r.source is not None + assert r.content.startswith(r.source) + + @pytest.mark.asyncio + async def test_a_closing_failure_does_not_mask_the_exit(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"]) + + rag = HaikuRAG(config=config) + await rag.__aenter__() + await rag.clients_for(["alpha", "beta"]) + assert isinstance(rag._session, FederatedSession) + sessions = rag._session._sessions + + async def boom(): + raise RuntimeError("close failed") + + sessions["alpha"].aclose = boom # ty: ignore[invalid-assignment] + beta = sessions["beta"].store + + await rag.__aexit__(None, None, None) + + # The failure is swallowed, and the sibling is still closed after it. + assert rag._clients == {} + assert rag._session._sessions == {} + assert not beta.db.is_open() + + @pytest.mark.asyncio + async def test_multimodal_reranking_attaches_each_database_own_pictures( + self, tmp_path, monkeypatch + ): + """Picture self_refs repeat across databases exactly as they do across + documents, so the pre-rerank attach must stay per database.""" + config = _config(tmp_path, ["alpha", "beta"]) + config.reranking.multimodal = True + dim = get_config().embeddings.model.vector_dim + + for name in ("alpha", "beta"): + async with HaikuRAG(config=config, create=True, sources=[name]) as rag: + doc = DoclingDocument(name=name) + doc.add_text(label=DocItemLabel.TEXT, text=f"{name} figure of cats") + await rag.import_document( + doc, + [ + Chunk( + content=f"{name} figure of cats", + embedding=[0.1] * dim, + order=0, + metadata={ + "doc_item_refs": ["#/pictures/0"], + "labels": ["picture"], + }, + ) + ], + uri=f"test://{name}/figure", + ) + [document] = await rag.list_documents() + assert document.id is not None + await rag.document_item_repository.create_items( + document.id, + [ + DocumentItem( + document_id=document.id, + position=0, + self_ref="#/pictures/0", + label="picture", + text=f"caption {name}", + picture_data=f"bytes-{name}".encode(), + ) + ], + ) + + stub = StubReranker() + monkeypatch.setattr(HaikuRAG, "reranker", property(lambda self: stub)) + + async with HaikuRAG(config=config) as rag: + await rag.search("cats", limit=2, search_type="fts") + + assert stub.attached == {"alpha": b"bytes-alpha", "beta": b"bytes-beta"} + + +class TestOneReranker: + @pytest.mark.asyncio + async def test_the_set_builds_one_reranker_for_a_text_query( + self, tmp_path, monkeypatch + ): + """Local rerankers load model weights per instance, so a set of + databases must build one, not one each.""" + built = [] + monkeypatch.setattr( + "haiku.rag.client.get_reranker", + lambda config: built.append(config) or StubReranker(), + ) + + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha document about cats"]) + await _seed(config, "beta", ["beta document about cats"]) + built.clear() + + async with HaikuRAG(config=config) as rag: + await rag.search("cats", limit=2, search_type="fts") + + assert len(built) == 1, f"built {len(built)} rerankers" + + @pytest.mark.asyncio + async def test_an_image_query_builds_no_reranker(self, tmp_path, monkeypatch): + """Opening a database must not build one either: an image query has no + text to score against and never uses it.""" + built = [] + monkeypatch.setattr( + "haiku.rag.client.get_reranker", + lambda config: built.append(config) or StubReranker(), + ) + + config = _config(tmp_path, ["alpha", "beta"]) + await _seed(config, "alpha", ["alpha document about cats"]) + await _seed(config, "beta", ["beta document about cats"]) + built.clear() + + async with HaikuRAG(config=config) as rag: + await rag.clients_for(["alpha", "beta"]) + + assert built == [] + + @pytest.mark.asyncio + async def test_the_reranker_is_closed_once(self, tmp_path, monkeypatch): + """Handing the same object to every database and letting each close it + would close it N times, and the federator not at all.""" + closes = [] + + class CountingReranker(StubReranker): + async def aclose(self): + closes.append(1) + + monkeypatch.setattr( + "haiku.rag.client.get_reranker", lambda config: CountingReranker() + ) + + 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: + await rag.search("cats", limit=2, search_type="fts") + + assert closes == [1], f"closed {len(closes)} times" + + +class TestOneNamedDatabase: + @pytest.mark.asyncio + async def test_a_single_named_database_keeps_its_name(self, tmp_path): + """Named in config is named in results, even as the only entry.""" + config = _config(tmp_path, ["alpha"]) + await _seed(config, "alpha", ["alpha document about cats"]) + + async with HaikuRAG(config=config) as rag: + results = await rag.search("cats", search_type="fts") + + assert results + assert all(r.source == "alpha" for r in results) + + @pytest.mark.asyncio + async def test_selecting_nothing_at_construction_is_rejected(self, tmp_path): + config = _config(tmp_path, ["alpha", "beta"]) + + with pytest.raises(ValueError, match="selects no database"): + async with HaikuRAG(config=config, sources=[]): + pass + + @pytest.mark.asyncio + async def test_selecting_nothing_means_the_same_with_one_database(self, tmp_path): + """`sources=[]` selects nothing whether one database is configured or + several, rather than raising on one path and returning nothing on the + other.""" + config = _config(tmp_path, ["alpha"]) + await _seed(config, "alpha", ["alpha document about cats"]) + + async with HaikuRAG(config=config) as rag: + assert await rag.search("cats", search_type="fts", sources=[]) == [] + + @pytest.mark.asyncio + async def test_selecting_nothing_per_query_returns_nothing(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: + assert await rag.search("cats", search_type="fts", sources=[]) == [] diff --git a/tests/sandbox/test_sandbox_multi_db.py b/tests/sandbox/test_sandbox_multi_db.py index 92c7c8cf..2ae69a7c 100644 --- a/tests/sandbox/test_sandbox_multi_db.py +++ b/tests/sandbox/test_sandbox_multi_db.py @@ -5,7 +5,7 @@ import pytest from haiku.rag.client import HaikuRAG from haiku.rag.client.scope import DatabaseRef, DatabaseScope from haiku.rag.sandbox import AnalysisContext, Sandbox -from tests.test_multi_db import _config, _seed +from tests.multi_db.helpers import _config, _seed async def _mounted(rag, sources=None): diff --git a/tests/test_multi_db.py b/tests/test_multi_db.py deleted file mode 100644 index 92365b41..00000000 --- a/tests/test_multi_db.py +++ /dev/null @@ -1,1405 +0,0 @@ -import asyncio - -import pytest -from docling_core.types.doc.document import DoclingDocument -from docling_core.types.doc.labels import DocItemLabel -from pydantic import ValidationError - -from haiku.rag.client import HaikuRAG -from haiku.rag.client.scope import DatabaseScope -from haiku.rag.client.session import FederatedSession -from haiku.rag.config import get_config -from haiku.rag.config.models import AppConfig, LanceDBConfig -from haiku.rag.store.exceptions import ( - AmbiguousDatabaseError, - ConfigMismatchError, - SourceUnavailableError, -) -from haiku.rag.store.models import Chunk, DocumentItem -from haiku.rag.utils import locate_database - - -class TestConfig: - def test_databases_and_uri_are_mutually_exclusive(self): - with pytest.raises(ValidationError, match="databases"): - LanceDBConfig( - uri="s3://b/one.lancedb", databases={"one": "s3://b/one.lancedb"} - ) - - def test_databases_alone_is_fine(self): - config = LanceDBConfig(databases={"one": "s3://b/one.lancedb"}) - assert config.databases == {"one": "s3://b/one.lancedb"} - - def test_uri_alone_is_fine(self): - assert LanceDBConfig(uri="s3://b/one.lancedb").databases == {} - - -class TestNamingIsRequired: - def test_a_blank_name_is_rejected(self): - """An unnamed database is unreachable: every source check reads the - empty name as no name at all.""" - with pytest.raises(ValidationError, match="entry with no name"): - LanceDBConfig(databases={"": "/tmp/a.lancedb"}) - with pytest.raises(ValidationError, match="entry with no name"): - LanceDBConfig(databases={" ": "/tmp/a.lancedb"}) - - def test_a_blank_location_is_rejected(self): - """A blank location resolves to the working directory.""" - with pytest.raises( - ValidationError, match=r"databases\[alpha\] has no location" - ): - LanceDBConfig(databases={"alpha": ""}) - - -def _config(tmp_path, names) -> AppConfig: - return AppConfig( - lancedb=LanceDBConfig( - databases={n: str(tmp_path / f"{n}.lancedb") for n in names} - ) - ) - - -async def _seed(config, name, contents): - """Precomputed embeddings and FTS queries keep the embedder out of the way: - these tests are about fusion, not retrieval quality.""" - dim = get_config().embeddings.model.vector_dim - async with HaikuRAG(config=config, create=True, sources=[name]) as rag: - for content in contents: - doc = DoclingDocument(name=content) - doc.add_text(label=DocItemLabel.TEXT, text=content) - await rag.import_document( - doc, - [Chunk(content=content, embedding=[0.1] * dim, order=0)], - uri=f"test://{name}/{content}", - ) - - -@pytest.fixture -def query_embedding(monkeypatch): - """Vector search with no embedder behind it, recording the queries embedded. - - These tests are about which databases are asked and how often, not about - retrieval quality, and CI has no embedding endpoint. - """ - from haiku.rag.embeddings import EmbedderWrapper - - embedded: list[str] = [] - - async def embed_query(self, text): - embedded.append(text) - return [0.1] * get_config().embeddings.model.vector_dim - - monkeypatch.setattr(EmbedderWrapper, "embed_query", embed_query) - return embedded - - -async def _restore_embedder(config, name, *, provider=None, model_name=None): - """Rewrite what one database records about the embedder that wrote it, - standing in for a database built elsewhere with another model.""" - import json - - import lancedb - - _, db_path = locate_database(config.lancedb.databases[name]) - assert db_path is not None - db = await lancedb.connect_async(str(db_path.resolve())) - table = await db.open_table("settings") - rows = ( - await table.query().where("id = 'settings'").limit(1).to_arrow() - ).to_pylist() - stored = json.loads(rows[0]["settings"]) - model = stored["embeddings"]["model"] - if provider is not None: - model["provider"] = provider - if model_name is not None: - model["name"] = model_name - await table.update({"settings": json.dumps(stored)}, where="id = 'settings'") - - -class TestNamingADatabaseDirectly: - @pytest.mark.asyncio - async def test_an_explicit_db_path_wins_over_the_configured_set( - self, tmp_path, temp_db_path - ): - """A caller that names a path means that database, not the configured - set: the CLI resolves `--db` to one and must not fan out instead.""" - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha document about cats"]) - - async with HaikuRAG(temp_db_path, config=config, create=True) as rag: - assert not rag.covers_multiple - assert rag.source is None - assert rag.store.db_path == temp_db_path - - @pytest.mark.asyncio - async def test_one_configured_database_is_opened_by_name(self, tmp_path): - """A set of one is not federated, and the client resolves it.""" - config = _config(tmp_path, ["alpha"]) - await _seed(config, "alpha", ["alpha document about cats"]) - - async with HaikuRAG(config=config) as rag: - assert not rag.covers_multiple - assert rag.source == "alpha" - results = await rag.search("cats", search_type="fts", limit=10) - - assert [r.source for r in results] == ["alpha"] - - -class TestOneConfiguredLocation: - """`lancedb.uri` places one unnamed database, at a URI or at a local path.""" - - def _config(self, location) -> AppConfig: - return AppConfig(lancedb=LanceDBConfig(uri=str(location))) - - @pytest.mark.asyncio - async def test_a_local_uri_opens_the_configured_database(self, tmp_path): - located = tmp_path / "notes.lancedb" - config = self._config(located) - - async with HaikuRAG(config=config, create=True) as rag: - assert rag.store.db_path == located - # It places a database without naming one: only `lancedb.databases` - # assigns the name results and citations carry. - assert rag.source is None - assert located.exists() - - @pytest.mark.asyncio - async def test_an_explicit_path_overrides_a_local_uri(self, tmp_path): - """`--db` overrides the configured location for one invocation.""" - config = self._config(tmp_path / "configured.lancedb") - chosen = tmp_path / "chosen.lancedb" - - async with HaikuRAG(chosen, config=config, create=True) as rag: - assert rag.store.db_path == chosen - assert chosen.exists() - assert not (tmp_path / "configured.lancedb").exists() - - @pytest.mark.asyncio - async def test_a_local_uri_that_does_not_exist_is_refused(self, tmp_path): - """A mistyped path fails instead of quietly becoming an empty database, - which is what a value carrying a scheme would do.""" - config = self._config(tmp_path / "typo.lancedb") - - with pytest.raises(FileNotFoundError): - async with HaikuRAG(config=config): - pass - assert not (tmp_path / "typo.lancedb").exists() - - def test_a_uri_with_a_scheme_stays_a_uri(self, tmp_path): - """Object storage has no local path to check, and a location that does - not exist yet is normal there.""" - from haiku.rag.store.engine import ConnectionMode - - config = self._config("s3://bucket/one.lancedb") - - [ref] = DatabaseScope.resolve(config).databases - one, db_path = ref.connection(config) - - assert db_path is None - assert ConnectionMode.from_config(one) == ConnectionMode.OBJECT_STORAGE - - -class TestOpeningDatabases: - @pytest.mark.asyncio - async def test_missing_databases_open_together(self, tmp_path): - """A cold fan-out costs one open, not their sum. On object storage a - serial loop is the difference between one round trip and N.""" - names = ["alpha", "beta", "gamma"] - config = _config(tmp_path, names) - for name in names: - await _seed(config, name, [f"{name} document about cats"]) - - async with HaikuRAG(config=config) as rag: - assert isinstance(rag._session, FederatedSession) - barrier = asyncio.Barrier(len(names)) - open_one = rag._session._open - - async def gated(ref): - # Every open has to be in flight before any of them finishes, so - # a serial loop cannot get past this and the wait times out. - await barrier.wait() - return await open_one(ref) - - rag._session._open = gated - clients = await asyncio.wait_for(rag.clients_for(names), timeout=15) - - assert {client.source for client in clients} == set(names) - - @pytest.mark.asyncio - async def test_a_failed_open_does_not_leak_the_ones_that_worked(self, tmp_path): - """Opening together means a failure has siblings already open. They are - tracked before it is reported, so closing the set closes them.""" - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha document about cats"]) - config.lancedb.databases["beta"] = str(tmp_path / "absent.lancedb") - - async with HaikuRAG(config=config) as rag: - with pytest.raises(SourceUnavailableError, match="beta"): - await rag.clients_for(["alpha", "beta"]) - - assert isinstance(rag._session, FederatedSession) - assert set(rag._session._sessions) == {"alpha"} - - @pytest.mark.asyncio - async def test_a_database_named_twice_is_opened_once(self, tmp_path): - """Fusion would count a repeated database as two rank lists.""" - 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: - clients = await rag.clients_for(["alpha", "alpha", "beta"]) - - assert [client.source for client in clients] == ["alpha", "beta"] - - @pytest.mark.asyncio - async def test_a_database_named_twice_returns_each_result_once(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: - results = await rag.search( - "cats", limit=10, search_type="fts", sources=["alpha", "alpha"] - ) - - assert [r.source for r in results] == ["alpha"] - - @pytest.mark.asyncio - async def test_one_database_named_twice_is_still_that_database(self, tmp_path): - """A client covering a single named database compares the selection - against its own name, so repeats have to collapse first.""" - config = _config(tmp_path, ["alpha"]) - await _seed(config, "alpha", ["alpha document about cats"]) - - async with HaikuRAG(config=config) as rag: - covering = await rag.clients_covering(["alpha", "alpha"]) - - assert [client.source for client in covering] == ["alpha"] - - -class TestListingAcrossDatabases: - """The chat TUI's document filter lists documents through the client, and a - client covering a set has no repositories of its own.""" - - @pytest.mark.asyncio - async def test_listing_covers_every_database(self, tmp_path): - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one", "alpha two"]) - await _seed(config, "beta", ["beta one"]) - - async with HaikuRAG(config=config) as rag: - docs = await rag.list_documents() - - assert {d.uri for d in docs} == { - "test://alpha/alpha one", - "test://alpha/alpha two", - "test://beta/beta one", - } - - @pytest.mark.asyncio - async def test_counting_covers_every_database(self, tmp_path): - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one", "alpha two"]) - await _seed(config, "beta", ["beta one"]) - - async with HaikuRAG(config=config) as rag: - assert await rag.count_documents() == 3 - - @pytest.mark.asyncio - async def test_a_limit_bounds_the_merged_listing(self, tmp_path): - """A limit is that many documents in total, not that many per database.""" - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one", "alpha two"]) - await _seed(config, "beta", ["beta one", "beta two"]) - - async with HaikuRAG(config=config) as rag: - assert len(await rag.list_documents(limit=3)) == 3 - assert len(await rag.list_documents(limit=2, offset=2)) == 2 - assert len(await rag.list_documents(offset=3)) == 1 - - @pytest.mark.asyncio - async def test_a_page_shows_every_database(self, tmp_path): - """A window is taken across the databases, not filled from the first one: - concatenating hides every database after whichever was listed first.""" - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", [f"alpha {i}" for i in range(5)]) - await _seed(config, "beta", ["beta one"]) - - async with HaikuRAG(config=config) as rag: - page = await rag.list_documents(limit=3) - - assert len(page) == 3 - assert {(d.uri or "").split("/")[2] for d in page} == {"alpha", "beta"} - - @pytest.mark.asyncio - async def test_a_filter_reaches_every_database(self, tmp_path): - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one"]) - await _seed(config, "beta", ["beta one"]) - - async with HaikuRAG(config=config) as rag: - docs = await rag.list_documents(filter="uri LIKE 'test://beta/%'") - - assert [d.uri for d in docs] == ["test://beta/beta one"] - - -class TestLookupByIdentifier: - """An id or a URI says nothing about which database holds it, and a client - covering a set has no repositories of its own.""" - - @pytest.mark.asyncio - async def test_a_document_is_found_in_whichever_database_holds_it(self, tmp_path): - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one"]) - await _seed(config, "beta", ["beta one"]) - - async with HaikuRAG(config=config) as rag: - beta = (await rag.clients_for(["beta"]))[0] - [target] = await beta.document_repository.list_all(limit=1) - assert target.id is not None - - found = await rag.get_document_by_id(target.id) - by_uri = await rag.get_document_by_uri("test://alpha/alpha one") - resolved = await rag.resolve_document(target.id) - - assert found is not None and found.uri == "test://beta/beta one" - assert by_uri is not None and by_uri.uri == "test://alpha/alpha one" - assert resolved is not None and resolved.uri == "test://beta/beta one" - - @pytest.mark.asyncio - async def test_a_chunk_is_found_in_whichever_database_holds_it(self, tmp_path): - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one"]) - await _seed(config, "beta", ["beta one"]) - - async with HaikuRAG(config=config) as rag: - beta = (await rag.clients_for(["beta"]))[0] - [chunk] = await beta.chunk_repository.list_all(limit=1) - assert chunk.id is not None - - found = await rag.get_chunk_by_id(chunk.id) - - assert found is not None and found.content == "beta one" - - @pytest.mark.asyncio - async def test_a_document_held_by_two_databases_answers_from_the_first( - self, tmp_path - ): - """A database copied from another holds the same ids. A read has an - answer wherever it finds one, and which one it is has to be the - configured order rather than whichever replied first.""" - import shutil - - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one"]) - shutil.copytree(tmp_path / "alpha.lancedb", tmp_path / "beta.lancedb") - - async with HaikuRAG(config=config) as rag: - beta = (await rag.clients_for(["beta"]))[0] - [target] = await beta.document_repository.list_all(limit=1) - assert target.id is not None - - found = await rag.get_document_by_id(target.id) - - assert found is not None and found.source == "alpha" - - @pytest.mark.asyncio - async def test_an_unknown_identifier_is_absent_rather_than_an_error(self, tmp_path): - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one"]) - await _seed(config, "beta", ["beta one"]) - - async with HaikuRAG(config=config) as rag: - assert ( - await rag.get_document_by_id("00000000-0000-4000-8000-000000000000") - is None - ) - assert ( - await rag.get_chunk_by_id("00000000-0000-4000-8000-000000000000") - is None - ) - assert await rag.get_document_by_uri("test://nowhere") is None - - -class TestClosingASet: - @pytest.mark.asyncio - async def test_every_database_opened_is_released(self, tmp_path): - """A covered database owns an embedder and may owe a vacuum. Closing only - its connection would leave both behind.""" - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one"]) - await _seed(config, "beta", ["beta one"]) - - released: list[str | None] = [] - drained: list[str | None] = [] - - async with HaikuRAG(config=config, read_only=True) as rag: - assert isinstance(rag._session, FederatedSession) - await rag.clients_for(["alpha", "beta"]) - for name, session in rag._session._sessions.items(): - original = session.store.embedder.aclose - drain = session.drain_vacuum - - async def release(_original=original, _name=name): - released.append(_name) - return await _original() - - async def drain_it(_drain=drain, _name=name): - drained.append(_name) - return await _drain() - - session.store.embedder.aclose = release - session.drain_vacuum = drain_it - - assert sorted(released) == ["alpha", "beta"] - assert sorted(drained) == ["alpha", "beta"] - - -class TestNamingOneOfTheSetOnTheCommandLine: - """`--database NAME` reaches the application layer as a name, and every - client it opens has to honour it — one that ignores it covers the set and - quietly answers from the wrong database.""" - - @pytest.mark.asyncio - async def test_a_named_database_is_the_one_read(self, tmp_path, capsys): - from haiku.rag.app import HaikuRAGApp - - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha document about cats"]) - await _seed(config, "beta", ["beta document about cats"]) - - scope = DatabaseScope.resolve(config).select(["beta"]) - app = HaikuRAGApp(scope=scope, config=config, read_only=True) - await app.list_documents() - - # Rich wraps long lines, so match the unwrapped part of the URI. - printed = capsys.readouterr().out - assert "test://beta/" in printed - assert "test://alpha/" not in printed - - @pytest.mark.asyncio - async def test_naming_none_of_them_covers_the_set(self, tmp_path, capsys): - from haiku.rag.app import HaikuRAGApp - - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha document about cats"]) - await _seed(config, "beta", ["beta document about cats"]) - - app = HaikuRAGApp( - scope=DatabaseScope.resolve(config), config=config, read_only=True - ) - await app.list_documents() - - printed = capsys.readouterr().out - assert "test://alpha/" in printed - assert "test://beta/" in printed - - -class TestPlacingADatabase: - """What a client says about the databases it covers, so nothing outside has - to read its private state to find out.""" - - @pytest.mark.asyncio - async def test_a_set_names_every_database_it_covers(self, tmp_path): - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one"]) - await _seed(config, "beta", ["beta one"]) - - async with HaikuRAG(config=config, read_only=True) as rag: - assert rag.covers_multiple - assert rag.source_names == ("alpha", "beta") - assert rag.source is None - - @pytest.mark.asyncio - async def test_one_named_database_names_itself(self, tmp_path): - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one"]) - - async with HaikuRAG(config=config, read_only=True, sources=["alpha"]) as rag: - assert not rag.covers_multiple - assert rag.source_names == ("alpha",) - assert rag.source == "alpha" - - @pytest.mark.asyncio - async def test_a_named_database_keeps_its_name_on_re_entry(self, tmp_path): - """Entering derives a single-database configuration from what was - configured. Deriving it from the last derivation loses the name.""" - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha document about cats"]) - - rag = HaikuRAG(config=config, read_only=True, sources=["alpha"]) - async with rag: - assert rag.source == "alpha" - async with rag: - assert rag.source == "alpha" - assert rag.source_names == ("alpha",) - results = await rag.search("cats", search_type="fts") - - assert {r.source for r in results} == {"alpha"} - - @pytest.mark.asyncio - async def test_an_unnamed_database_names_nothing(self, temp_db_path): - async with HaikuRAG(temp_db_path, create=True) as rag: - assert rag.source_names == () - assert rag.source is None - - @pytest.mark.asyncio - async def test_the_reader_for_a_database_is_the_client_holding_it(self, tmp_path): - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one"]) - await _seed(config, "beta", ["beta one"]) - - async with HaikuRAG(config=config, read_only=True) as rag: - reader = await rag.reader_for("beta") - - assert reader is not None - assert reader.source == "beta" - # Asked twice, the same wrapper comes back. - assert await rag.reader_for("beta") is reader - - @pytest.mark.asyncio - async def test_a_client_reading_one_database_is_its_own_reader(self, temp_db_path): - async with HaikuRAG(temp_db_path, create=True) as rag: - assert await rag.reader_for(None) is rag - assert await rag.reader_for("anything") is rag - - @pytest.mark.asyncio - async def test_a_set_cannot_place_evidence_that_names_no_database(self, tmp_path): - """Evidence recorded before databases could be named carries no source.""" - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one"]) - - async with HaikuRAG(config=config, read_only=True) as rag: - assert await rag.reader_for(None) is None - - -class TestBorrowedDatabases: - """A client for one of a set wraps a database the set opened.""" - - @pytest.mark.asyncio - async def test_closing_a_borrowed_client_leaves_the_set_working(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: - (alpha,) = await rag.clients_for(["alpha"]) - store = alpha.store - - alpha.close() - assert store.db.is_open(), "close() closed a database it borrowed" - - await alpha.__aexit__(None, None, None) - assert store.db.is_open(), "exit closed a database it borrowed" - - results = await rag.search("cats", search_type="fts") - - assert {r.source for r in results} == {"alpha", "beta"} - assert not store.db.is_open(), "the set left a database open" - - @pytest.mark.asyncio - async def test_entering_a_borrowed_client_reuses_its_database(self, tmp_path): - """`async with` on a borrowed client is a plausible thing to write. - Opening a second session would leak it, since teardown declines to close - what this client did not open.""" - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha document about cats"]) - - async with HaikuRAG(config=config) as rag: - (alpha,) = await rag.clients_for(["alpha"]) - borrowed = alpha.store - - async with alpha as entered: - assert entered is alpha - assert alpha.store is borrowed, "entry opened a second database" - - assert borrowed.db.is_open(), "exit closed a database it borrowed" - assert alpha.store is borrowed - - assert not borrowed.db.is_open(), "the set left a database open" - - @pytest.mark.asyncio - async def test_a_borrowed_client_releases_what_it_built(self, tmp_path): - """Its reranker is its own; the database it wraps is not.""" - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha document about cats"]) - - closed: list[str] = [] - - class Reranker: - async def aclose(self): - closed.append("reranker") - - async with HaikuRAG(config=config) as rag: - (alpha,) = await rag.clients_for(["alpha"]) - alpha.__dict__["reranker"] = Reranker() - - assert closed == ["reranker"] - - -class TestDatabaseIndependentWork: - """Converting, chunking and titling are functions of the configuration, not - of a database, so covering a set does not stop them.""" - - @pytest.mark.asyncio - async def test_chunking_opens_no_database(self, tmp_path, monkeypatch): - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one"]) - await _seed(config, "beta", ["beta one"]) - - opened: list[str] = [] - - async def refuse(self, ref): - opened.append(ref.name) - raise AssertionError("opened a database to chunk a document") - - monkeypatch.setattr(FederatedSession, "_open", refuse) - - doc = DoclingDocument(name="note") - doc.add_text( - label=DocItemLabel.TEXT, text="Boltzmann machines are energy based." - ) - - async with HaikuRAG(config=config, read_only=True) as rag: - chunks = await rag.chunk(doc) - - assert opened == [] - assert [c.content for c in chunks] - - @pytest.mark.asyncio - async def test_the_embedder_is_built_once_and_closed_once( - self, tmp_path, monkeypatch - ): - """The parent owns the embedder it built, so leaving the context closes - it, once.""" - from haiku.rag.embeddings import EmbedderWrapper - - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one"]) - - closed: list[object] = [] - original = EmbedderWrapper.aclose - - async def counting(self): - closed.append(self) - return await original(self) - - monkeypatch.setattr(EmbedderWrapper, "aclose", counting) - - rag = HaikuRAG(config=config, read_only=True) - async with rag: - built = rag.embedder - assert rag.embedder is built - - assert closed == [built] - - @pytest.mark.asyncio - async def test_re_entering_a_set_builds_a_fresh_embedder(self, tmp_path): - """Teardown closes the embedder, so keeping it would hand the next - context one that is already closed.""" - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one"]) - - rag = HaikuRAG(config=config, read_only=True) - async with rag: - first = rag.embedder - async with rag: - assert rag.embedder is not first - - @pytest.mark.asyncio - async def test_re_entering_one_database_builds_a_fresh_embedder(self, temp_db_path): - """One database opens a new store on re-entry, and the embedder is that - store's.""" - rag = HaikuRAG(temp_db_path, create=True) - async with rag: - first = rag.embedder - async with rag: - assert rag.embedder is rag.store.embedder - assert rag.embedder is not first - - @pytest.mark.asyncio - async def test_a_set_nobody_asked_anything_of_builds_no_embedder(self, tmp_path): - """Built on first use, so a client that answered nothing holds nothing.""" - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one"]) - - async with HaikuRAG(config=config, read_only=True) as rag: - assert "embedder" not in rag.__dict__ - - @pytest.mark.asyncio - async def test_one_database_still_uses_its_store_s_embedder(self, temp_db_path): - async with HaikuRAG(temp_db_path, create=True) as rag: - assert rag.embedder is rag.store.embedder - - -class TestCreatingNeedsOneDatabase: - """Creating names a database. Covering a set, the flag had nothing to act on - and was accepted anyway, leaving the first query to fail on whichever - database turned out to be missing.""" - - @pytest.mark.asyncio - async def test_creating_a_set_is_refused(self, tmp_path): - config = _config(tmp_path, ["alpha", "beta"]) - - with pytest.raises(AmbiguousDatabaseError, match="alpha, beta"): - async with HaikuRAG(config=config, create=True): - pass - - @pytest.mark.asyncio - async def test_naming_one_of_the_set_creates_it(self, tmp_path): - config = _config(tmp_path, ["alpha", "beta"]) - - async with HaikuRAG(config=config, create=True, sources=["alpha"]) as rag: - assert await rag.count_documents() == 0 - - assert (tmp_path / "alpha.lancedb").exists() - assert not (tmp_path / "beta.lancedb").exists() - - @pytest.mark.asyncio - async def test_covering_a_set_without_creating_is_unaffected(self, tmp_path): - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one"]) - await _seed(config, "beta", ["beta one"]) - - async with HaikuRAG(config=config) as rag: - assert await rag.count_documents() == 2 - - -class TestOperationsThatNeedOneDatabase: - @pytest.mark.asyncio - async def test_writing_names_the_databases_it_covers(self, tmp_path): - """A domain error, so a caller can tell an unsupported selection from a - missing attribute.""" - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one"]) - await _seed(config, "beta", ["beta one"]) - - async with HaikuRAG(config=config) as rag: - with pytest.raises(AmbiguousDatabaseError, match="alpha, beta"): - await rag.create_document("orphan") - with pytest.raises(AmbiguousDatabaseError, match="clients_for"): - await rag.vacuum() - with pytest.raises(AmbiguousDatabaseError, match="close"): - rag.close() - - @pytest.mark.asyncio - async def test_a_set_has_no_store_of_its_own(self, tmp_path): - """A store and its repositories belong to one database. `clients_for` - reaches the one holding a given database.""" - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one"]) - - async with HaikuRAG(config=config) as rag: - for name in ( - "store", - "document_repository", - "chunk_repository", - "document_item_repository", - ): - with pytest.raises(AttributeError, match=name): - getattr(rag, name) - - @pytest.mark.asyncio - async def test_a_selected_database_is_still_writable(self, tmp_path): - """Naming one of the set is how a write picks its database.""" - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one"]) - await _seed(config, "beta", ["beta one"]) - - dim = get_config().embeddings.model.vector_dim - written = DoclingDocument(name="written") - written.add_text(label=DocItemLabel.TEXT, text="written") - - async with HaikuRAG(config=config) as rag: - alpha = (await rag.clients_for(["alpha"]))[0] - assert alpha.is_read_only is False - document = await alpha.import_document( - written, - [Chunk(content="written", embedding=[0.1] * dim, order=0)], - uri="test://alpha/written", - ) - assert await alpha.count_documents() == 2 - - assert document.id is not None - - -class TestOneQueryVector: - @pytest.mark.asyncio - async def test_a_search_embeds_the_query_once_for_the_whole_set( - self, tmp_path, query_embedding - ): - """Each database owns an embedder, so embedding per database costs a - round trip each on a remote endpoint.""" - config = _config(tmp_path, ["alpha", "beta", "gamma"]) - for name in ("alpha", "beta", "gamma"): - await _seed(config, name, [f"{name} one"]) - - async with HaikuRAG(config=config, read_only=True) as rag: - await rag.search("one") - - assert query_embedding == ["one"] - - -class TestOneEmbedderAcrossTheSet: - """A set is searched with one query vector, so a database written with - another model would answer from a different space.""" - - @pytest.mark.asyncio - async def test_disagreeing_databases_cannot_be_searched_together(self, tmp_path): - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one"]) - await _seed(config, "beta", ["beta one"]) - await _restore_embedder(config, "beta", model_name="some-other-model") - - async with HaikuRAG(config=config, read_only=True) as rag: - with pytest.raises(ConfigMismatchError, match="different embedders"): - await rag.search("one") - - @pytest.mark.asyncio - async def test_a_database_asked_for_alone_is_never_compared( - self, tmp_path, query_embedding - ): - """Only databases searched together have to agree.""" - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one"]) - await _seed(config, "beta", ["beta one"]) - await _restore_embedder(config, "beta", model_name="some-other-model") - - async with HaikuRAG(config=config, read_only=True) as rag: - assert await rag.search("one", sources=["alpha"]) is not None - assert await rag.count_documents(filter=None) is not None - - @pytest.mark.asyncio - async def test_full_text_search_needs_no_agreement(self, tmp_path): - """Full-text search embeds nothing, so which model wrote each database - does not come into it.""" - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one"]) - await _seed(config, "beta", ["beta one"]) - await _restore_embedder(config, "beta", model_name="some-other-model") - - async with HaikuRAG(config=config, read_only=True) as rag: - results = await rag.search("one", search_type="fts") - - assert {r.source for r in results} == {"alpha", "beta"} - - @pytest.mark.asyncio - async def test_agreeing_databases_search_together(self, tmp_path, query_embedding): - """The databases agree with each other; that they were written by a - differently-spelled provider than the config is the soft case.""" - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one"]) - await _seed(config, "beta", ["beta one"]) - await _restore_embedder(config, "alpha", provider="openai") - await _restore_embedder(config, "beta", provider="openai") - - async with HaikuRAG(config=config, read_only=True) as rag: - assert len(await rag.search("one")) > 0 - - -class TestReadOnlyMode: - @pytest.mark.asyncio - async def test_a_client_covering_a_set_reports_its_mode(self, tmp_path): - """A client covering a set has no store of its own to ask.""" - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one"]) - await _seed(config, "beta", ["beta one"]) - - async with HaikuRAG(config=config, read_only=True) as rag: - assert rag.is_read_only is True - async with HaikuRAG(config=config) as rag: - assert rag.is_read_only is False - - -class TestDocumentsNameTheirDatabase: - """A listing that spans databases is unreadable when the documents do not - say which one they came from, the same reason a search result carries one.""" - - @pytest.mark.asyncio - async def test_a_listing_names_each_document_s_database(self, tmp_path): - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one", "alpha two"]) - await _seed(config, "beta", ["beta one"]) - - async with HaikuRAG(config=config) as rag: - docs = await rag.list_documents() - - assert {d.uri: d.source for d in docs} == { - "test://alpha/alpha one": "alpha", - "test://alpha/alpha two": "alpha", - "test://beta/beta one": "beta", - } - - @pytest.mark.asyncio - async def test_a_looked_up_document_names_its_database(self, tmp_path): - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one"]) - await _seed(config, "beta", ["beta one"]) - - async with HaikuRAG(config=config) as rag: - beta = (await rag.clients_for(["beta"]))[0] - [target] = await beta.document_repository.list_all(limit=1) - assert target.id is not None - - by_id = await rag.get_document_by_id(target.id) - by_uri = await rag.get_document_by_uri("test://alpha/alpha one") - resolved = await rag.resolve_document(target.id) - - assert by_id is not None and by_id.source == "beta" - assert by_uri is not None and by_uri.source == "alpha" - assert resolved is not None and resolved.source == "beta" - - @pytest.mark.asyncio - async def test_one_named_database_still_names_itself(self, tmp_path): - """`haiku-rag --database alpha list` opens one database, and its name is - the whole reason the option exists.""" - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one"]) - await _seed(config, "beta", ["beta one"]) - - async with HaikuRAG(config=config, sources=["alpha"]) as rag: - [listed] = await rag.list_documents() - assert listed.id is not None - by_id = await rag.get_document_by_id(listed.id) - by_uri = await rag.get_document_by_uri("test://alpha/alpha one") - - assert listed.source == "alpha" - assert by_id is not None and by_id.source == "alpha" - assert by_uri is not None and by_uri.source == "alpha" - - @pytest.mark.asyncio - async def test_one_database_leaves_the_source_unset(self, tmp_path, temp_db_path): - """Nothing names the database when there is only one to name.""" - async with HaikuRAG(temp_db_path, create=True) as rag: - dim = get_config().embeddings.model.vector_dim - doc = DoclingDocument(name="solo") - doc.add_text(label=DocItemLabel.TEXT, text="solo") - await rag.import_document( - doc, - [Chunk(content="solo", embedding=[0.1] * dim, order=0)], - uri="test://solo", - ) - - [listed] = await rag.list_documents() - assert listed.source is None - assert listed.id is not None - assert (await rag.get_document_by_id(listed.id)).source is None - - -class TestFederatedSearch: - @pytest.mark.asyncio - async def test_results_carry_their_source(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: - results = await rag.search("cats", limit=10, search_type="fts") - - assert {r.source for r in results} == {"alpha", "beta"} - for r in results: - assert r.source is not None - assert r.source in r.content - - @pytest.mark.asyncio - async def test_sources_selects_a_subset(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: - results = await rag.search( - "cats", limit=10, search_type="fts", sources=["alpha"] - ) - - assert {r.source for r in results} == {"alpha"} - - @pytest.mark.asyncio - async def test_unknown_source_is_rejected(self, tmp_path): - config = _config(tmp_path, ["alpha"]) - await _seed(config, "alpha", ["alpha document about cats"]) - - async with HaikuRAG(config=config) as rag: - with pytest.raises(KeyError, match="nope"): - await rag.search("cats", search_type="fts", sources=["nope"]) - - @pytest.mark.asyncio - async def test_an_unopenable_database_fails_the_query(self, tmp_path): - config = _config(tmp_path, ["alpha", "missing"]) - await _seed(config, "alpha", ["alpha document about cats"]) - - with pytest.raises(SourceUnavailableError, match="missing"): - async with HaikuRAG(config=config) as rag: - await rag.search("cats", search_type="fts") - - -class TestSingleDatabaseUnchanged: - @pytest.mark.asyncio - async def test_source_is_unset_without_configured_databases(self, temp_db_path): - async with HaikuRAG(temp_db_path, create=True) as rag: - doc = DoclingDocument(name="one") - doc.add_text(label=DocItemLabel.TEXT, text="a document about cats") - await rag.import_document( - doc, - [ - Chunk( - content="a document about cats", - embedding=[0.1] * get_config().embeddings.model.vector_dim, - order=0, - ) - ], - uri="test://one", - ) - results = await rag.search("cats", search_type="fts") - - assert results - assert all(r.source is None for r in results) - - -class TestLocate: - def test_a_scheme_is_a_uri(self): - assert locate_database("s3://bucket/one.lancedb") == ( - "s3://bucket/one.lancedb", - None, - ) - - def test_anything_else_is_a_local_path(self): - uri, db_path = locate_database("/data/one.lancedb") - assert uri == "" - assert db_path is not None and str(db_path) == "/data/one.lancedb" - - -class TestSelection: - @pytest.mark.asyncio - async def test_unknown_source_at_construction_is_rejected(self, tmp_path): - config = _config(tmp_path, ["alpha", "beta"]) - - with pytest.raises(KeyError, match="nope"): - async with HaikuRAG(config=config, sources=["nope"]): - pass - - @pytest.mark.asyncio - async def test_unknown_source_across_several_databases_is_rejected(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="nope"): - await rag.search("cats", search_type="fts", sources=["nope"]) - - @pytest.mark.asyncio - async def test_no_matches_anywhere_returns_nothing(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: - assert await rag.search("aardvarks", search_type="fts") == [] - - -class StubReranker: - """Scores the union, reversing it so the ordering is unmistakably its own.""" - - def __init__(self): - self.seen: list[str] = [] - - async def rerank(self, query, chunks, top_n): - self.seen = [c.content for c in chunks] - # Whatever the caller attached before handing them over. - self.attached = { - c.content.split()[0]: c._picture_data - for c in chunks - if getattr(c, "_picture_data", None) - } - return [(c, 1.0 - i) for i, c in enumerate(reversed(chunks))][:top_n] - - -class TestRerankerFusion: - @pytest.mark.asyncio - async def test_the_reranker_scores_the_union_and_owners_survive( - self, tmp_path, monkeypatch - ): - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha document about cats"]) - await _seed(config, "beta", ["beta document about cats"]) - - stub = StubReranker() - monkeypatch.setattr(HaikuRAG, "reranker", property(lambda self: stub)) - - async with HaikuRAG(config=config) as rag: - results = await rag.search("cats", limit=2, search_type="fts") - - # It saw both databases' candidates, not one database at a time. - assert len(stub.seen) == 2 - assert {c.split()[0] for c in stub.seen} == {"alpha", "beta"} - # Each result still knows which database it came from. - for r in results: - assert r.source is not None - assert r.content.startswith(r.source) - - @pytest.mark.asyncio - async def test_a_closing_failure_does_not_mask_the_exit(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"]) - - rag = HaikuRAG(config=config) - await rag.__aenter__() - await rag.clients_for(["alpha", "beta"]) - assert isinstance(rag._session, FederatedSession) - sessions = rag._session._sessions - - async def boom(): - raise RuntimeError("close failed") - - sessions["alpha"].aclose = boom # ty: ignore[invalid-assignment] - beta = sessions["beta"].store - - await rag.__aexit__(None, None, None) - - # The failure is swallowed, and the sibling is still closed after it. - assert rag._clients == {} - assert rag._session._sessions == {} - assert not beta.db.is_open() - - @pytest.mark.asyncio - async def test_multimodal_reranking_attaches_each_database_own_pictures( - self, tmp_path, monkeypatch - ): - """Picture self_refs repeat across databases exactly as they do across - documents, so the pre-rerank attach must stay per database.""" - config = _config(tmp_path, ["alpha", "beta"]) - config.reranking.multimodal = True - dim = get_config().embeddings.model.vector_dim - - for name in ("alpha", "beta"): - async with HaikuRAG(config=config, create=True, sources=[name]) as rag: - doc = DoclingDocument(name=name) - doc.add_text(label=DocItemLabel.TEXT, text=f"{name} figure of cats") - await rag.import_document( - doc, - [ - Chunk( - content=f"{name} figure of cats", - embedding=[0.1] * dim, - order=0, - metadata={ - "doc_item_refs": ["#/pictures/0"], - "labels": ["picture"], - }, - ) - ], - uri=f"test://{name}/figure", - ) - [document] = await rag.list_documents() - assert document.id is not None - await rag.document_item_repository.create_items( - document.id, - [ - DocumentItem( - document_id=document.id, - position=0, - self_ref="#/pictures/0", - label="picture", - text=f"caption {name}", - picture_data=f"bytes-{name}".encode(), - ) - ], - ) - - stub = StubReranker() - monkeypatch.setattr(HaikuRAG, "reranker", property(lambda self: stub)) - - async with HaikuRAG(config=config) as rag: - await rag.search("cats", limit=2, search_type="fts") - - assert stub.attached == {"alpha": b"bytes-alpha", "beta": b"bytes-beta"} - - -class TestLazyOpening: - @pytest.mark.asyncio - async def test_entering_opens_nothing(self, tmp_path): - """25 configured databases queried a few at a time must not all open.""" - 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 rag._clients == {} - - @pytest.mark.asyncio - async def test_only_the_selected_database_opens(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: - await rag.search("cats", search_type="fts", sources=["alpha"]) - assert list(rag._clients) == ["alpha"] - - @pytest.mark.asyncio - async def test_an_unselected_broken_database_does_not_break_a_query(self, tmp_path): - """A database nobody asked for cannot fail a query.""" - config = _config(tmp_path, ["alpha", "missing"]) - await _seed(config, "alpha", ["alpha document about cats"]) - - async with HaikuRAG(config=config) as rag: - results = await rag.search("cats", search_type="fts", sources=["alpha"]) - - assert [r.source for r in results] == ["alpha"] - - @pytest.mark.asyncio - async def test_a_selected_broken_database_fails_the_query(self, tmp_path): - config = _config(tmp_path, ["alpha", "missing"]) - await _seed(config, "alpha", ["alpha document about cats"]) - - async with HaikuRAG(config=config) as rag: - with pytest.raises(SourceUnavailableError, match="missing"): - await rag.search("cats", search_type="fts") - - -class TestOneNamedDatabase: - @pytest.mark.asyncio - async def test_a_single_named_database_keeps_its_name(self, tmp_path): - """Named in config is named in results, even as the only entry.""" - config = _config(tmp_path, ["alpha"]) - await _seed(config, "alpha", ["alpha document about cats"]) - - async with HaikuRAG(config=config) as rag: - results = await rag.search("cats", search_type="fts") - - assert results - assert all(r.source == "alpha" for r in results) - - @pytest.mark.asyncio - async def test_selecting_nothing_at_construction_is_rejected(self, tmp_path): - config = _config(tmp_path, ["alpha", "beta"]) - - with pytest.raises(ValueError, match="selects no database"): - async with HaikuRAG(config=config, sources=[]): - pass - - @pytest.mark.asyncio - async def test_selecting_nothing_means_the_same_with_one_database(self, tmp_path): - """`sources=[]` selects nothing whether one database is configured or - several, rather than raising on one path and returning nothing on the - other.""" - config = _config(tmp_path, ["alpha"]) - await _seed(config, "alpha", ["alpha document about cats"]) - - async with HaikuRAG(config=config) as rag: - assert await rag.search("cats", search_type="fts", sources=[]) == [] - - @pytest.mark.asyncio - async def test_selecting_nothing_per_query_returns_nothing(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: - assert await rag.search("cats", search_type="fts", sources=[]) == [] - - -class TestOneReranker: - @pytest.mark.asyncio - async def test_the_set_builds_one_reranker_for_a_text_query( - self, tmp_path, monkeypatch - ): - """Local rerankers load model weights per instance, so a set of - databases must build one, not one each.""" - built = [] - monkeypatch.setattr( - "haiku.rag.client.get_reranker", - lambda config: built.append(config) or StubReranker(), - ) - - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha document about cats"]) - await _seed(config, "beta", ["beta document about cats"]) - built.clear() - - async with HaikuRAG(config=config) as rag: - await rag.search("cats", limit=2, search_type="fts") - - assert len(built) == 1, f"built {len(built)} rerankers" - - @pytest.mark.asyncio - async def test_an_image_query_builds_no_reranker(self, tmp_path, monkeypatch): - """Opening a database must not build one either: an image query has no - text to score against and never uses it.""" - built = [] - monkeypatch.setattr( - "haiku.rag.client.get_reranker", - lambda config: built.append(config) or StubReranker(), - ) - - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha document about cats"]) - await _seed(config, "beta", ["beta document about cats"]) - built.clear() - - async with HaikuRAG(config=config) as rag: - await rag.clients_for(["alpha", "beta"]) - - assert built == [] - - @pytest.mark.asyncio - async def test_the_reranker_is_closed_once(self, tmp_path, monkeypatch): - """Handing the same object to every database and letting each close it - would close it N times, and the federator not at all.""" - closes = [] - - class CountingReranker(StubReranker): - async def aclose(self): - closes.append(1) - - monkeypatch.setattr( - "haiku.rag.client.get_reranker", lambda config: CountingReranker() - ) - - 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: - await rag.search("cats", limit=2, search_type="fts") - - assert closes == [1], f"closed {len(closes)} times" - - -class TestFailureNaming: - @pytest.mark.asyncio - async def test_a_single_named_database_is_reported_by_name(self, tmp_path): - """One configured database is still a named one: it must not fall back to - the raw error, which spells out the path.""" - config = _config(tmp_path, ["alpha"]) - - with pytest.raises(SourceUnavailableError, match="alpha") as caught: - async with HaikuRAG(config=config): - pass - - assert str(tmp_path) not in str(caught.value) - assert caught.value.__cause__ is None - - @pytest.mark.asyncio - async def test_a_legacy_uri_client_keeps_its_error(self, tmp_path): - """Nothing named it, so there is no name to report instead.""" - with pytest.raises(FileNotFoundError): - async with HaikuRAG(tmp_path / "nope.lancedb"): - pass - - @pytest.mark.asyncio - async def test_the_location_is_absent_from_the_whole_chain(self, tmp_path): - config = _config(tmp_path, ["alpha", "missing"]) - await _seed(config, "alpha", ["alpha document about cats"]) - - with pytest.raises(SourceUnavailableError) as caught: - async with HaikuRAG(config=config) as rag: - await rag.search("cats", search_type="fts") - - rendered = str(caught.value) - error = caught.value.__cause__ or caught.value.__context__ - assert "missing.lancedb" not in rendered - assert error is None, "the location-bearing cause is still attached" diff --git a/tests/test_multi_db_ask.py b/tests/test_multi_db_ask.py deleted file mode 100644 index c0aa321d..00000000 --- a/tests/test_multi_db_ask.py +++ /dev/null @@ -1,814 +0,0 @@ -import pytest -from docling_core.types.doc.document import DoclingDocument -from docling_core.types.doc.labels import DocItemLabel -from pydantic_ai import ModelRetry - -from haiku.rag.capabilities.rag import RAGState, create_capability -from haiku.rag.client import HaikuRAG -from haiku.rag.config import get_config -from haiku.rag.store.exceptions import AmbiguousCitationError -from haiku.rag.store.models import Chunk, Document, DocumentItem, SearchResult -from haiku.rag.store.models.citation import Citation, resolve_citations -from tests.test_multi_db import _config, _seed - - -async def _seed_expandable(config, name, sentences): - """One document whose chunk covers a single item, so expansion has - neighbours to pull in and rebuilds the result rather than passing it - through.""" - dim = get_config().embeddings.model.vector_dim - doc = DoclingDocument(name=name) - for sentence in sentences: - doc.add_text(label=DocItemLabel.TEXT, text=sentence) - async with HaikuRAG(config=config, create=True, sources=[name]) as rag: - await rag.import_document( - doc, - [ - Chunk( - content=sentences[0], - embedding=[0.1] * dim, - order=0, - metadata={"doc_item_refs": ["#/texts/0"]}, - ) - ], - uri=f"test://{name}/expandable", - ) - - -class TestExpansionRouting: - @pytest.mark.asyncio - async def test_expansion_routes_each_result_to_its_database(self, tmp_path): - """A federating client has no repositories of its own, so expansion has - to go through the database each result came from.""" - 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: - results = await rag.search("cats", search_type="fts", limit=10) - expanded = await rag.expand_context(results) - - assert {r.source for r in expanded} == {"alpha", "beta"} - for r in expanded: - assert r.source is not None - assert r.source in r.content - - @pytest.mark.asyncio - async def test_an_expanded_result_keeps_its_source(self, tmp_path): - """Expansion rebuilds the result, and the rebuilt one has to name the - database it was expanded through.""" - config = _config(tmp_path, ["alpha"]) - await _seed_expandable( - config, "alpha", ["cats sleep often", "cats also hunt", "cats purr"] - ) - - async with HaikuRAG(config=config) as rag: - results = await rag.search("cats", search_type="fts", limit=10) - expanded = await rag.expand_context(results) - - assert len(expanded) == 1 - assert "cats also hunt" in expanded[0].content, "expansion did not run" - assert expanded[0].source == "alpha" - - @pytest.mark.asyncio - async def test_a_federated_result_is_expanded_by_its_own_database(self, tmp_path): - """Routing is not enough: each result has to come back carrying the - neighbours of the database it was expanded through, and only those.""" - config = _config(tmp_path, ["alpha", "beta"]) - await _seed_expandable( - config, "alpha", ["cats sleep often", "alpha follows on"] - ) - await _seed_expandable(config, "beta", ["cats also hunt", "beta follows on"]) - - async with HaikuRAG(config=config) as rag: - results = await rag.search("cats", search_type="fts", limit=10) - expanded = await rag.expand_context(results) - - content = {r.source: r.content for r in expanded} - assert "alpha follows on" in content["alpha"] - assert "beta follows on" not in content["alpha"] - assert "beta follows on" in content["beta"] - - @pytest.mark.asyncio - async def test_expansion_keeps_tied_results_in_fused_order(self, tmp_path): - """Fused scores tie often, so grouping by database must not reorder - them: the tiebreak is the order they arrived in.""" - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one about cats", "alpha two about cats"]) - await _seed(config, "beta", ["beta one about cats"]) - - async with HaikuRAG(config=config) as rag: - found = await rag.search("cats", search_type="fts", limit=10) - by_source: dict[str, list[SearchResult]] = {} - for result in found: - by_source.setdefault(result.source or "", []).append(result) - # Interleaved, so grouping by database is visible as a reordering. - fused = [by_source["alpha"][0], by_source["beta"][0], by_source["alpha"][1]] - for result in fused: - result.score = 0.5 - - expanded = await rag.expand_context(fused) - - assert [r.chunk_id for r in expanded] == [r.chunk_id for r in fused] - - -class TestSharedChunkIds: - """A database copied from another holds the same chunk ids.""" - - @pytest.mark.asyncio - async def test_a_shared_id_does_not_confuse_the_fused_order(self, tmp_path): - """Arrival order breaks score ties, so it has to tell two databases' - identically-numbered chunks apart.""" - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha one about cats"]) - await _seed(config, "beta", ["beta one about cats"]) - fused = [ - SearchResult(content="a0", score=0.5, chunk_id="a0", source="alpha"), - SearchResult(content="beta", score=0.5, chunk_id="shared", source="beta"), - SearchResult(content="alpha", score=0.5, chunk_id="shared", source="alpha"), - ] - - async with HaikuRAG(config=config) as rag: - expanded = await rag.expand_context(fused) - - assert [(r.source, r.chunk_id) for r in expanded] == [ - (r.source, r.chunk_id) for r in fused - ] - - def test_a_shared_id_cannot_be_cited(self): - """A citation records the id alone, so resolving one held by two - databases would attribute the answer to whichever came last.""" - results = [ - SearchResult( - content="alpha body", - score=0.9, - source="alpha", - chunk_id="c1", - document_id="d1", - document_uri="test://alpha/one", - ), - SearchResult( - content="beta body", - score=0.8, - source="beta", - chunk_id="c1", - document_id="d1", - document_uri="test://beta/one", - ), - ] - - with pytest.raises(AmbiguousCitationError, match="c1"): - resolve_citations(["c1"], results) - - def test_a_repeated_id_from_one_database_still_collapses(self): - """One database cannot hold two chunks under one id, so seeing it twice - is the same chunk seen twice.""" - results = [ - SearchResult( - content="first", - score=0.9, - source="alpha", - chunk_id="c1", - document_id="d1", - document_uri="test://alpha/one", - ), - SearchResult( - content="second", - score=0.8, - source="alpha", - chunk_id="c1", - document_id="d1", - document_uri="test://alpha/one", - ), - ] - - [citation] = resolve_citations(["c1"], results) - - assert citation.content == "first" - - def test_only_a_cited_id_has_to_be_unambiguous(self): - """An id the answer never cites attributes nothing.""" - shared = [ - SearchResult( - content=f"{name} body", - score=0.9, - source=name, - chunk_id="c1", - document_id="d1", - document_uri=f"test://{name}/one", - ) - for name in ("alpha", "beta") - ] - own = SearchResult( - content="alpha only", - score=0.7, - source="alpha", - chunk_id="c2", - document_id="d2", - document_uri="test://alpha/two", - ) - - [citation] = resolve_citations(["c2"], [*shared, own]) - - assert citation.source == "alpha" - - @pytest.mark.asyncio - async def test_an_unsearched_shared_id_is_refused_by_the_fallback(self, tmp_path): - """The direct lookup is the only place a collision shows for an id no - search returned, so it has to ask every database rather than take the - first that answers.""" - import shutil - - config = _config(tmp_path, ["alpha", "beta"]) - await _seed( - config, "alpha", ["alpha document about cats", "alpha on aardvarks"] - ) - shutil.copytree(tmp_path / "alpha.lancedb", tmp_path / "beta.lancedb") - - async with HaikuRAG(config=config) as rag: - alpha = (await rag.clients_for(["alpha"]))[0] - chunks = await alpha.chunk_repository.list_all() - [aardvark] = [c for c in chunks if "aardvark" in c.content] - assert aardvark.id is not None - - capability = create_capability(config=config, rag=rag, defer_loading=False) - capability.state = RAGState() - - # No search ran, so the id can only resolve through the fallback. - with pytest.raises(ModelRetry, match="more than one database"): - await capability._cite([aardvark.id]) - - @pytest.mark.asyncio - async def test_an_unsearched_id_in_one_database_still_resolves(self, tmp_path): - """The refusal is for a collision, not for looking through several - databases: an id only one of them holds still resolves.""" - from tests.capabilities.test_capabilities import Deps, make_context - - config = _config(tmp_path, ["alpha", "beta"]) - await _seed( - config, "alpha", ["alpha document about cats", "alpha on aardvarks"] - ) - await _seed(config, "beta", ["beta document about dogs"]) - - async with HaikuRAG(config=config) as rag: - alpha = (await rag.clients_for(["alpha"]))[0] - chunks = await alpha.chunk_repository.list_all() - [aardvark] = [c for c in chunks if "aardvark" in c.content] - assert aardvark.id is not None - - capability = create_capability(config=config, rag=rag, defer_loading=False) - run = await capability.for_run(make_context(Deps())) - - await run._cite([aardvark.id]) - - assert run.state is not None - [citation] = list(run.state.citation_index.values()) - assert citation.source == "alpha" - - @pytest.mark.asyncio - async def test_cite_asks_for_other_evidence(self, tmp_path): - capability = create_capability( - config=_config(tmp_path, ["alpha", "beta"]), defer_loading=False - ) - capability.state = RAGState( - searches={ - "cats": [ - SearchResult( - content=f"{name} body", - score=0.9, - source=name, - chunk_id="c1", - document_id="d1", - document_uri=f"test://{name}/one", - ) - for name in ("alpha", "beta") - ] - } - ) - - with pytest.raises(ModelRetry, match="appears once"): - await capability._cite(["c1"]) - - @pytest.mark.asyncio - async def test_cite_refuses_an_id_already_cited_from_another_database( - self, tmp_path - ): - """The citation index outlives the question, so the collision can arrive - a turn later than the search that would have shown it.""" - capability = create_capability( - config=_config(tmp_path, ["alpha", "beta"]), defer_loading=False - ) - capability.state = RAGState( - citation_index={ - "c1": Citation( - document_id="d1", - source="alpha", - chunk_id="c1", - document_uri="test://alpha/one", - content="alpha body", - ) - }, - searches={ - "cats": [ - SearchResult( - content="beta body", - score=0.9, - source="beta", - chunk_id="c1", - document_id="d1", - document_uri="test://beta/one", - ) - ] - }, - ) - - with pytest.raises(ModelRetry, match="another database"): - await capability._cite(["c1"]) - - -class TestCitationSource: - def test_a_citation_carries_the_result_source(self): - result = SearchResult( - content="body", - score=0.9, - source="alpha", - chunk_id="c1", - document_id="d1", - document_uri="test://alpha/one", - ) - - [citation] = resolve_citations(["c1"], [result]) - - assert citation.source == "alpha" - - def test_a_result_without_an_id_is_skipped(self): - """A result built by hand carries no chunk id, so nothing can cite it - and it takes part in no collision.""" - handmade = SearchResult(content="loose text", score=0.5) - real = SearchResult( - content="body", - score=0.9, - source="alpha", - chunk_id="c1", - document_id="d1", - document_uri="test://alpha/one", - ) - - [citation] = resolve_citations(["c1"], [handmade, real]) - - assert citation.chunk_id == "c1" - - def test_a_single_database_citation_has_no_source(self): - result = SearchResult( - content="body", - score=0.9, - chunk_id="c1", - document_id="d1", - document_uri="test://one", - ) - - [citation] = resolve_citations(["c1"], [result]) - - assert citation.source is None - - -class TestAskAcrossDatabases: - @pytest.mark.asyncio - @pytest.mark.vcr() - async def test_the_capability_searches_the_selected_databases(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: - capability = create_capability(config=config, rag=rag, defer_loading=False) - capability.state = RAGState(sources=["alpha"]) - - formatted = await capability._search("cats", limit=10) - - assert isinstance(formatted, str) - assert "alpha" in formatted - assert "beta document" not in formatted - - @pytest.mark.asyncio - @pytest.mark.vcr() - async def test_searching_all_databases_reaches_both(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: - capability = create_capability(config=config, rag=rag, defer_loading=False) - capability.state = RAGState() - - formatted = await capability._search("cats", limit=10) - - assert isinstance(formatted, str) - assert "alpha document" in formatted - assert "beta document" in formatted - - -class TestCiteFallback: - @pytest.mark.asyncio - @pytest.mark.vcr() - async def test_an_id_from_a_selected_database_resolves_with_its_source( - self, tmp_path - ): - """The fallback exists for a real id this run's searches did not return. - Across databases it looks through the selected ones and records which - held it.""" - from tests.capabilities.test_capabilities import Deps, make_context - - config = _config(tmp_path, ["alpha", "beta"]) - await _seed( - config, "alpha", ["alpha document about cats", "alpha on aardvarks"] - ) - await _seed(config, "beta", ["beta document about cats"]) - - async with HaikuRAG(config=config) as rag: - alpha = (await rag.clients_for(["alpha"]))[0] - chunks = await alpha.chunk_repository.list_all() - [aardvark] = [c for c in chunks if "aardvark" in c.content] - assert aardvark.id is not None - - capability = create_capability(config=config, rag=rag, defer_loading=False) - deps = Deps( - state={"rag": RAGState(sources=["alpha"]).model_dump(mode="json")} - ) - run = await capability.for_run(make_context(deps)) - # The search returns the cats chunk, never the aardvark one. - await run._search("cats", limit=10) - - await run._cite([aardvark.id]) - - assert run.state is not None - [citation] = list(run.state.citation_index.values()) - assert citation.chunk_id == aardvark.id - assert citation.source == "alpha" - - @pytest.mark.asyncio - @pytest.mark.vcr() - async def test_an_id_outside_the_selected_databases_does_not_resolve( - self, tmp_path - ): - """A question scoped to one database must not produce a citation from - another: the fallback looks only where the question looked.""" - from tests.capabilities.test_capabilities import Deps, make_context - - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha document about cats"]) - await _seed(config, "beta", ["beta document about dogs"]) - - async with HaikuRAG(config=config) as rag: - beta = (await rag.clients_for(["beta"]))[0] - [outside] = await beta.chunk_repository.list_all(limit=1) - assert outside.id is not None - - capability = create_capability(config=config, rag=rag, defer_loading=False) - deps = Deps( - state={"rag": RAGState(sources=["alpha"]).model_dump(mode="json")} - ) - run = await capability.for_run(make_context(deps)) - await run._search("cats", limit=10) - - with pytest.raises(ModelRetry): - await run._cite([outside.id]) - - @pytest.mark.asyncio - async def test_selecting_no_databases_cites_nothing(self, tmp_path): - """`sources=[]` selected nothing, which is not the same as everything: - the fallback must not go looking where the question never looked.""" - from tests.capabilities.test_capabilities import Deps, make_context - - 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: - alpha = (await rag.clients_for(["alpha"]))[0] - [chunk] = await alpha.chunk_repository.list_all(limit=1) - assert chunk.id is not None - - capability = create_capability(config=config, rag=rag, defer_loading=False) - deps = Deps(state={"rag": RAGState(sources=[]).model_dump(mode="json")}) - run = await capability.for_run(make_context(deps)) - - with pytest.raises(ModelRetry): - await run._cite([chunk.id]) - - -class TestStandaloneCapabilities: - """A capability nobody hands a client opens its own. It has to reach the - configured set, or a host that only registers capabilities gets one - database while the configuration names several.""" - - @pytest.mark.asyncio - @pytest.mark.vcr() - async def test_a_rag_capability_opens_the_configured_set(self, tmp_path): - from tests.capabilities.test_capabilities import Deps, make_context - - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha document about cats"]) - await _seed(config, "beta", ["beta document about cats"]) - - capability = create_capability(config=config, defer_loading=False) - assert capability.scope.names == ("alpha", "beta") - run = await capability.for_run(make_context(Deps())) - try: - formatted = await run._search("cats", limit=10) - finally: - await run._close() - - assert isinstance(formatted, str) - assert "alpha document" in formatted - assert "beta document" in formatted - - @pytest.mark.asyncio - @pytest.mark.vcr() - async def test_an_analysis_capability_mounts_the_configured_set(self, tmp_path): - from haiku.rag.capabilities.analysis import ( - create_capability as create_analysis, - ) - from tests.capabilities.test_capabilities import Deps, make_context - - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha document about cats"]) - await _seed(config, "beta", ["beta document about cats"]) - - capability = create_analysis(config=config, defer_loading=False) - run = await capability.for_run(make_context(Deps())) - try: - sandbox = await run._ensure_sandbox() - docs, owners = await sandbox._documents() - finally: - await run._close() - - assert len(docs) == 2 - assert {owner.source for owner in owners.values()} == {"alpha", "beta"} - - @pytest.mark.asyncio - async def test_a_single_configured_database_is_still_opened(self, tmp_path): - """One named database is a set of one, not a path to guess.""" - config = _config(tmp_path, ["alpha"]) - await _seed(config, "alpha", ["alpha document about cats"]) - - capability = create_capability(config=config, defer_loading=False) - rag = await capability._ensure_rag() - try: - assert rag.source == "alpha" - finally: - await capability._close() - - -class TestAnalyzeAcrossDatabases: - @pytest.mark.asyncio - @pytest.mark.vcr() - async def test_the_capability_searches_the_selected_databases(self, tmp_path): - """`analysis_search` is the same tool as the RAG one, and the sandbox is - scoped by the same selection.""" - from haiku.rag.capabilities.analysis import AnalysisState - from haiku.rag.capabilities.analysis import ( - create_capability as create_analysis, - ) - - 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: - capability = create_analysis(config=config, rag=rag, defer_loading=False) - capability.state = AnalysisState(sources=["alpha"]) - - formatted = await capability._search("cats", limit=10) - sandbox = await capability._ensure_sandbox() - await capability._close() - - assert isinstance(formatted, str) - assert "alpha document" in formatted - assert "beta document" not in formatted - assert sandbox._context.sources == ["alpha"] - - -class TestDatabaseIdentityForTheModel: - def test_a_result_names_its_database(self): - """The model has to attribute and compare evidence by database while it - composes the answer, not only afterwards through the citations.""" - result = SearchResult(content="body", score=0.9, source="alpha", chunk_id="c1") - - assert "Database: alpha" in result.format_for_agent() - - def test_an_unnamed_database_is_not_mentioned(self): - """A single unnamed database renders as it always has.""" - result = SearchResult(content="body", score=0.9, chunk_id="c1") - - assert "Database" not in result.format_for_agent() - - @pytest.mark.asyncio - @pytest.mark.vcr() - async def test_in_code_search_names_the_database(self, tmp_path): - from haiku.rag.sandbox import AnalysisContext, Sandbox - - 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: - sandbox = Sandbox( - db_path=None, - config=config, - context=AnalysisContext(), - rag=rag, - ) - try: - result = await sandbox.execute( - "rows = await search('cats', limit=10)\n" - "print(sorted(r['source'] for r in rows))\n" - "docs = await list_documents()\n" - "print(sorted(d['source'] for d in docs))" - ) - finally: - await sandbox.close() - - assert result.success, result.stderr - assert "['alpha', 'beta']" in result.stdout - assert result.stdout.count("['alpha', 'beta']") == 2 - - -class TestActionableFailures: - @pytest.mark.asyncio - async def test_a_migration_error_survives_being_named(self, tmp_path, temp_db_path): - """The remedy is the whole value of the message, and it names no location, - so it is not replaced by the database's name.""" - from haiku.rag.store.exceptions import MigrationRequiredError - - config = _config(tmp_path, ["alpha"]) - await _seed(config, "alpha", ["alpha document about cats"]) - - async with HaikuRAG(config=config, sources=["alpha"]) as rag: - await rag.store.set_haiku_version("0.20.0") - - with pytest.raises(MigrationRequiredError) as raised: - async with HaikuRAG(config=config, sources=["alpha"]): - pass - - # Both halves: which database failed, and what to run about it. - assert "haiku-rag migrate" in str(raised.value) - assert "alpha" in str(raised.value) - assert str(tmp_path) not in str(raised.value) - - -class TestPictureDeduplication: - """One picture yields two chunks — a text-embedded one and an image-embedded - one — that collapse to the best. Two databases holding the same picture are - two results, not a duplicate.""" - - @staticmethod - def _picture(source, score): - return SearchResult( - content="a figure", - score=score, - source=source, - chunk_id=f"{source}-c", - document_id="doc-1", - doc_item_refs=["#/pictures/0"], - ) - - def test_the_same_picture_in_two_databases_survives(self): - from haiku.rag.client.search import _dedup_picture_chunks - - kept = _dedup_picture_chunks( - [self._picture("alpha", 0.9), self._picture("clone", 0.5)] - ) - - assert [r.source for r in kept] == ["alpha", "clone"] - - def test_duplicates_within_one_database_still_collapse(self): - from haiku.rag.client.search import _dedup_picture_chunks - - lower = self._picture("alpha", 0.5) - higher = self._picture("alpha", 0.9) - - kept = _dedup_picture_chunks([lower, higher]) - - assert kept == [higher] - - -class TestPictureRouting: - @pytest.mark.asyncio - async def test_a_picture_is_fetched_from_the_database_that_holds_it(self, tmp_path): - """A `self_ref` repeats across databases, so the citation's source is - what decides where the bytes come from.""" - 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: - beta = (await rag.clients_for(["beta"]))[0] - [document] = await beta.document_repository.list_all(limit=1) - assert document.id is not None - await beta.document_item_repository.create_all( - [ - DocumentItem( - document_id=document.id, - self_ref="#/pictures/0", - position=99, - label="picture", - text="", - picture_data=b"beta-picture", - ) - ] - ) - - assert ( - await rag.get_picture_bytes(document.id, "#/pictures/0", "beta") - == b"beta-picture" - ) - assert ( - await rag.get_picture_bytes(document.id, "#/pictures/0", "alpha") - is None - ) - - @pytest.mark.asyncio - async def test_a_single_database_needs_no_source(self, temp_db_path): - """One database is where the picture is, named or not.""" - async with HaikuRAG(temp_db_path, create=True) as rag: - document = await rag.document_repository.create( - Document(content="body", uri="test://one") - ) - assert document.id is not None - await rag.document_item_repository.create_all( - [ - DocumentItem( - document_id=document.id, - self_ref="#/pictures/0", - position=0, - label="picture", - text="", - picture_data=b"the-picture", - ) - ] - ) - - assert ( - await rag.get_picture_bytes(document.id, "#/pictures/0") - == b"the-picture" - ) - - @pytest.mark.asyncio - async def test_a_picture_lookup_without_a_source_is_refused(self, tmp_path): - """Federating, nothing can say which database holds an unqualified - reference.""" - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha document about cats"]) - - async with HaikuRAG(config=config) as rag: - with pytest.raises(ValueError, match="source"): - await rag.get_picture_bytes("doc-1", "#/pictures/0") - - -class TestFederatedEdges: - @pytest.mark.asyncio - async def test_expansion_passes_through_results_without_a_source(self, tmp_path): - """A caller can hand `expand_context` results it built itself. Those name - no database, so there is nowhere to expand them from.""" - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha document about cats"]) - await _seed(config, "beta", ["beta document about cats"]) - - handmade = SearchResult(content="handmade", score=0.4, doc_item_refs=[]) - - async with HaikuRAG(config=config) as rag: - found = await rag.search("cats", search_type="fts", limit=10) - expanded = await rag.expand_context([*found, handmade]) - - assert "handmade" in [r.content for r in expanded] - scores = [r.score for r in expanded] - assert scores == sorted(scores, reverse=True), "merged in score order" - - @pytest.mark.asyncio - async def test_a_chunk_without_a_document_is_not_cited(self, tmp_path): - """`Chunk.document_id` is optional, and a citation without a document has - nothing to point at.""" - from unittest.mock import AsyncMock, patch - - from haiku.rag.capabilities.rag import RAGCapability - from haiku.rag.store.models import Chunk - from tests.capabilities.test_capabilities import ( - Deps, - _single_database_client, - make_context, - ) - - config = _config(tmp_path, ["alpha", "beta"]) - await _seed(config, "alpha", ["alpha document about cats"]) - - orphan = _single_database_client() - orphan.get_chunk_by_id.return_value = Chunk( - id="orphan", document_id=None, content="no document" - ) - - capability = create_capability(config=config, defer_loading=False) - run = await capability.for_run(make_context(Deps())) - with patch.object(RAGCapability, "_ensure_rag", AsyncMock(return_value=orphan)): - with pytest.raises(ModelRetry): - await run._cite(["orphan"])