haiku.rag/tests/sandbox/test_sandbox_multi_db.py
Yiorgis Gozadinos afdef92b5b
Finish the comment pass, and escape document fields everywhere Rich renders
`_rich_print_document` escapes uri, title and metadata, the sibling of
the escaped search-result renderer. The remaining comments and
docstrings that narrated rejected alternatives, consequences or history
now state the current invariant. The Sandbox class docstring names the
held connection close() releases, and wrapped docs paragraphs join to
one line.
2026-08-28 15:34:47 +03:00

403 lines
16 KiB
Python

import shutil
import pytest
from haiku.rag.client import HaikuRAG
from haiku.rag.client.scope import DatabaseRef, DatabaseScope
from haiku.rag.sandbox import AnalysisContext, Sandbox
from haiku.rag.store.exceptions import UnknownDatabaseError
from tests.multi_db.helpers import _config, _seed
async def _mounted(rag, sources=None):
"""The sandbox's view of the corpus, and the sandbox itself."""
sandbox = Sandbox(
db_path=None,
config=rag._config,
context=AnalysisContext(sources=sources),
rag=rag,
)
docs, owners = await sandbox._documents()
return sandbox, docs, owners
class TestSerializingTheConnection:
"""The lock guards the shared connection, which the capability's own tool
calls also hold. An owner is a session of its own."""
@staticmethod
def _sandbox(rag, lock):
return Sandbox._covering(
rag._resolve_scope(), rag._config, AnalysisContext(), rag, lock
)
@pytest.mark.asyncio
async def test_the_shared_connection_is_serialized(self, tmp_path):
import asyncio
config = _config(tmp_path, ["alpha", "beta"])
await _seed(config, "alpha", ["alpha document about cats"])
lock = asyncio.Lock()
async with HaikuRAG(config=config) as rag:
sandbox = self._sandbox(rag, lock)
async with sandbox._connection():
assert lock.locked()
assert not lock.locked()
@pytest.mark.asyncio
async def test_an_owner_is_not(self, tmp_path):
"""The lock guards the lent session's state, which owner reads do not
touch: they take no lock."""
import asyncio
class Trap(asyncio.Lock):
"""Raises on acquire, failing the test at the serialization
point."""
async def acquire(self):
raise AssertionError("serialized a read on an owner's own session")
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"])
sandbox = self._sandbox(rag, Trap())
async with sandbox._connection(alpha) as connection:
assert connection is alpha
class TestStandaloneAcrossDatabases:
"""Without a lent client the sandbox opens its own. The owners it hands
out are stored for later file reads and outlive the call that produced
them."""
@pytest.mark.asyncio
async def test_owners_stay_open_for_later_reads(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"])
sandbox = Sandbox(db_path=None, config=config, context=AnalysisContext())
try:
_, owners = await sandbox._documents()
assert len(owners) == 2
assert all(owner.store.db.is_open() for owner in owners.values())
finally:
await sandbox.close()
assert not any(owner.store.db.is_open() for owner in owners.values())
class TestDocumentsAcrossDatabases:
@pytest.mark.asyncio
async def test_the_corpus_covers_every_configured_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:
_, docs, owners = await _mounted(rag)
assert {d.uri for d in docs} == {
"test://alpha/alpha document about cats",
"test://beta/beta document about cats",
}
assert {owner.source for owner in owners.values()} == {"alpha", "beta"}
@pytest.mark.asyncio
async def test_selected_databases_bound_the_corpus(self, tmp_path):
"""A question scoped to one database must not mount another's documents."""
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:
_, docs, owners = await _mounted(rag, sources=["alpha"])
assert [d.uri for d in docs] == ["test://alpha/alpha document about cats"]
assert {owner.source for owner in owners.values()} == {"alpha"}
@pytest.mark.asyncio
async def test_one_database_needs_no_owners(self, tmp_path, temp_db_path):
"""A single connection serves every read, so nothing has to be routed."""
config = _config(tmp_path, ["alpha"])
await _seed(config, "alpha", ["alpha document about cats"])
async with HaikuRAG(config=config) as rag:
_, docs, owners = await _mounted(rag)
assert len(docs) == 1
assert owners == {}
@pytest.mark.asyncio
async def test_a_document_is_read_from_the_database_holding_it(self, tmp_path):
"""Reads addressed to one document go through its owner, which is the
only client that can answer them."""
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, docs, owners = await _mounted(rag)
sandbox._owners = owners
for doc in docs:
assert doc.id is not None
async with sandbox._connection(owners[doc.id]) as owner:
content = await owner.document_repository.get_content(doc.id)
assert content is not None
assert owners[doc.id].source is not None
assert owners[doc.id].source in content
class TestTheSandboxConstructors:
"""`Sandbox` is public and takes a path; `_covering` is for callers that
already resolved a scope, as `HaikuRAG._covering` is."""
@pytest.mark.asyncio
async def test_the_public_constructor_resolves_the_path_it_is_given(self, tmp_path):
config = _config(tmp_path, ["alpha", "beta"])
await _seed(config, "alpha", ["alpha document about cats"])
sandbox = Sandbox(
db_path=tmp_path / "alpha.lancedb",
config=config,
context=AnalysisContext(),
)
assert sandbox._scope.databases == (DatabaseRef.at(tmp_path / "alpha.lancedb"),)
@pytest.mark.asyncio
async def test_no_path_covers_what_the_configuration_places(self, tmp_path):
config = _config(tmp_path, ["alpha", "beta"])
await _seed(config, "alpha", ["alpha document about cats"])
sandbox = Sandbox(db_path=None, config=config, context=AnalysisContext())
assert sandbox._scope.names == ("alpha", "beta")
@pytest.mark.asyncio
async def test_covering_resolves_nothing_of_its_own(self, tmp_path, monkeypatch):
"""Handed a scope, it resolves nothing of its own."""
config = _config(tmp_path, ["alpha", "beta"])
await _seed(config, "alpha", ["alpha document about cats"])
scope = DatabaseScope.resolve(config, database_name="alpha")
def _refuse(*args, **kwargs):
raise AssertionError("resolved a scope it was already given")
monkeypatch.setattr(DatabaseScope, "resolve", _refuse)
sandbox = Sandbox._covering(scope, config, AnalysisContext())
assert sandbox._scope is scope
assert sandbox._config is config
class TestTheSandboxCoversWhatTheCapabilityCovers:
@pytest.mark.asyncio
async def test_the_capability_hands_over_the_scope_it_resolved(self, tmp_path):
"""The sandbox covers the scope the capability resolved, as handed
over."""
from haiku.rag.capabilities.analysis import AnalysisState, create_capability
config = _config(tmp_path, ["alpha", "beta"])
await _seed(config, "alpha", ["alpha document about cats"])
capability = create_capability(
db_path=tmp_path / "alpha.lancedb", config=config, defer_loading=False
)
capability.state = AnalysisState()
sandbox = await capability._ensure_sandbox()
try:
assert sandbox._scope is capability.scope
assert capability.scope.names == ()
finally:
await capability._close()
class TestExecutingAcrossDatabases:
@pytest.mark.asyncio
async def test_code_reads_documents_from_every_database(self, tmp_path):
"""The virtual filesystem is one flat namespace over the whole selected
set, so code reads a document without knowing which database holds it."""
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=rag._config,
context=AnalysisContext(),
rag=rag,
)
try:
result = await sandbox.execute(
"docs = await list_documents()\n"
"for d in sorted(docs, key=lambda d: d['uri']):\n"
" with open('/documents/' + d['id'] + '/content.txt') as f:\n"
" print(f.read())"
)
finally:
await sandbox.close()
assert result.success, result.stderr
assert "alpha document about cats" in result.stdout
assert "beta document about cats" in result.stdout
@pytest.mark.asyncio
async def test_code_cannot_read_an_unselected_database(self, tmp_path):
"""Scoping the question scopes the filesystem."""
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]
[outside] = await beta.document_repository.list_all(limit=1)
sandbox = Sandbox(
db_path=None,
config=rag._config,
context=AnalysisContext(sources=["alpha"]),
rag=rag,
)
try:
result = await sandbox.execute(
"docs = await list_documents()\n"
"print(len(docs))\n"
f"print(open('/documents/{outside.id}/content.txt').read())"
)
finally:
await sandbox.close()
assert not result.success
assert "beta document" not in result.stdout
class TestSelectionOnOneDatabase:
"""A client covering a single named database answers a selection the same way
a search does: the sandbox mounts what a search reaches."""
@pytest.mark.asyncio
async def test_selecting_no_database_mounts_nothing(self, tmp_path):
config = _config(tmp_path, ["alpha"])
await _seed(config, "alpha", ["alpha document about cats"])
async with HaikuRAG(config=config) as rag:
_, docs, owners = await _mounted(rag, sources=[])
assert docs == []
assert owners == {}
@pytest.mark.asyncio
async def test_selecting_another_database_is_refused(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(UnknownDatabaseError, match="beta"):
await _mounted(rag, sources=["beta"])
@pytest.mark.asyncio
async def test_selecting_it_by_name_mounts_it(self, tmp_path):
config = _config(tmp_path, ["alpha"])
await _seed(config, "alpha", ["alpha document about cats"])
async with HaikuRAG(config=config) as rag:
_, docs, _ = await _mounted(rag, sources=["alpha"])
assert len(docs) == 1
class TestCopiedDatabases:
@pytest.mark.asyncio
async def test_a_document_in_two_databases_is_refused(self, tmp_path):
"""Ids are unique per database, not across a copy of one, so one path
cannot serve two documents."""
config = _config(tmp_path, ["alpha", "clone"])
await _seed(config, "alpha", ["alpha document about cats"])
shutil.rmtree(tmp_path / "clone.lancedb", ignore_errors=True)
shutil.copytree(tmp_path / "alpha.lancedb", tmp_path / "clone.lancedb")
async with HaikuRAG(config=config) as rag:
with pytest.raises(ValueError, match="one document per id"):
await _mounted(rag)
@pytest.mark.asyncio
async def test_the_refusal_names_the_databases(self, tmp_path):
config = _config(tmp_path, ["alpha", "clone"])
await _seed(config, "alpha", ["alpha document about cats"])
shutil.rmtree(tmp_path / "clone.lancedb", ignore_errors=True)
shutil.copytree(tmp_path / "alpha.lancedb", tmp_path / "clone.lancedb")
async with HaikuRAG(config=config) as rag:
with pytest.raises(ValueError) as raised:
await _mounted(rag)
assert "alpha" in str(raised.value)
assert "clone" in str(raised.value)
assert str(tmp_path) not in str(raised.value)
class TestListingOrder:
@pytest.mark.asyncio
async def test_the_listing_interleaves_the_databases(self, tmp_path):
"""A truncated listing still shows documents from every database."""
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:
sandbox = Sandbox(
db_path=None,
config=config,
context=AnalysisContext(),
rag=rag,
)
docs, _ = await sandbox._documents()
assert len(docs) == 6
# The head has to reveal both databases.
assert {(d.uri or "").split("/")[2] for d in docs[:2]} == {"alpha", "beta"}
@pytest.mark.asyncio
async def test_in_code_list_documents_names_the_database(self, tmp_path):
"""`source` is what lets code group the corpus by database."""
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:
sandbox = Sandbox(
db_path=None,
config=config,
context=AnalysisContext(),
rag=rag,
)
rows = await sandbox._build_external_functions()["list_documents"]()
assert "source" in rows[0]
assert {r["source"] for r in rows} == {"alpha", "beta"}
@pytest.mark.asyncio
async def test_in_code_list_documents_names_one_database_too(self, tmp_path):
"""A document knows which database it came from whether or not the
analysis spans several."""
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:
sandbox = Sandbox(
db_path=None,
config=config,
context=AnalysisContext(),
rag=rag,
)
rows = await sandbox._build_external_functions()["list_documents"]()
assert [r["source"] for r in rows] == ["alpha"]