Find the owner of an id through one primitive
`first_found` replaces the client's private version and the capability's sequential one. The info modal reports through the connection the client already holds.
This commit is contained in:
parent
307e250f29
commit
f35be76604
4 changed files with 75 additions and 36 deletions
|
|
@ -33,9 +33,9 @@ from haiku.rag.capabilities._tools import (
|
||||||
search_corpus,
|
search_corpus,
|
||||||
)
|
)
|
||||||
from haiku.rag.capabilities.ledger import CapabilityEvidenceRecord, EvidenceRef
|
from haiku.rag.capabilities.ledger import CapabilityEvidenceRecord, EvidenceRef
|
||||||
from haiku.rag.client import HaikuRAG
|
from haiku.rag.client import HaikuRAG, first_found
|
||||||
from haiku.rag.config.models import AppConfig
|
from haiku.rag.config.models import AppConfig
|
||||||
from haiku.rag.store.models.chunk import Chunk, SearchResult
|
from haiku.rag.store.models.chunk import SearchResult
|
||||||
from haiku.rag.store.models.citation import Citation, resolve_citations
|
from haiku.rag.store.models.citation import Citation, resolve_citations
|
||||||
from haiku.rag.tools.search import build_image_content_from_results
|
from haiku.rag.tools.search import build_image_content_from_results
|
||||||
|
|
||||||
|
|
@ -159,21 +159,6 @@ def _called_own_tool(messages: list[ModelMessage], tool_names: frozenset[str]) -
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
async def _first_holding(
|
|
||||||
clients: "list[HaikuRAG]", chunk_id: str
|
|
||||||
) -> "tuple[HaikuRAG, Chunk] | None":
|
|
||||||
"""The first client holding this chunk, and the chunk.
|
|
||||||
|
|
||||||
A chunk id says nothing about which database holds it, so the only way to
|
|
||||||
place one is to ask. Returns None when none of them has it.
|
|
||||||
"""
|
|
||||||
for client in clients:
|
|
||||||
chunk = await client.get_chunk_by_id(chunk_id)
|
|
||||||
if chunk is not None:
|
|
||||||
return client, chunk
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class RAGCapabilityBase[StateT: EvidenceState](AbstractCapability[Any]):
|
class RAGCapabilityBase[StateT: EvidenceState](AbstractCapability[Any]):
|
||||||
db_path: Path | None
|
db_path: Path | None
|
||||||
|
|
@ -550,7 +535,9 @@ class RAGCapabilityBase[StateT: EvidenceState](AbstractCapability[Any]):
|
||||||
synthetic: list[SearchResult] = []
|
synthetic: list[SearchResult] = []
|
||||||
documents: dict[tuple[str | None, str], Any] = {}
|
documents: dict[tuple[str | None, str], Any] = {}
|
||||||
for chunk_id in missing:
|
for chunk_id in missing:
|
||||||
found = await _first_holding(lookups, chunk_id)
|
found = await first_found(
|
||||||
|
lookups, lambda owner: owner.get_chunk_by_id(chunk_id)
|
||||||
|
)
|
||||||
if found is None:
|
if found is None:
|
||||||
continue
|
continue
|
||||||
owner, chunk = found
|
owner, chunk = found
|
||||||
|
|
|
||||||
|
|
@ -60,6 +60,24 @@ _VACUUM_MIN_INTERVAL_S = 300.0
|
||||||
_NAMEABLE_FAILURES = (MigrationRequiredError, ConfigMismatchError, ReadOnlyError)
|
_NAMEABLE_FAILURES = (MigrationRequiredError, ConfigMismatchError, ReadOnlyError)
|
||||||
|
|
||||||
|
|
||||||
|
async def first_found(
|
||||||
|
clients: "list[HaikuRAG]",
|
||||||
|
lookup: "Callable[[HaikuRAG], Coroutine[Any, Any, Any]]",
|
||||||
|
) -> "tuple[HaikuRAG, Any] | None":
|
||||||
|
"""The first of `clients` for which `lookup` finds something, and what it found.
|
||||||
|
|
||||||
|
An id or a URI says nothing about which database holds it, so every one is
|
||||||
|
asked at once and the first that has it, in the order given, answers. Asking
|
||||||
|
in turn would cost a round trip per database for an identifier that is
|
||||||
|
missing or held by the last of them.
|
||||||
|
"""
|
||||||
|
found_by_client = await asyncio.gather(*(lookup(client) for client in clients))
|
||||||
|
for client, found in zip(clients, found_by_client, strict=True):
|
||||||
|
if found is not None:
|
||||||
|
return client, found
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
def _spell(embedding: tuple[str | None, str | None, int | None]) -> str:
|
def _spell(embedding: tuple[str | None, str | None, int | None]) -> str:
|
||||||
"""An embedder identity, for an error message."""
|
"""An embedder identity, for an error message."""
|
||||||
provider, name, vector_dim = embedding
|
provider, name, vector_dim = embedding
|
||||||
|
|
@ -748,18 +766,9 @@ class HaikuRAG:
|
||||||
async def _from_any_covered(
|
async def _from_any_covered(
|
||||||
self, lookup: "Callable[[HaikuRAG], Coroutine[Any, Any, Any]]"
|
self, lookup: "Callable[[HaikuRAG], Coroutine[Any, Any, Any]]"
|
||||||
) -> Any:
|
) -> Any:
|
||||||
"""The first result `lookup` finds in the databases this client covers.
|
"""The first result `lookup` finds in the databases this client covers."""
|
||||||
|
found = await first_found(await self.clients_covering(), lookup)
|
||||||
An id or a URI says nothing about which database holds it, so every
|
return None if found is None else found[1]
|
||||||
database is asked at once and the first that has it, in configured order,
|
|
||||||
answers. Asking in turn would cost a round trip per database for an
|
|
||||||
identifier that is missing or held by the last of them.
|
|
||||||
"""
|
|
||||||
owners = await self.clients_covering()
|
|
||||||
for found in await asyncio.gather(*(lookup(owner) for owner in owners)):
|
|
||||||
if found is not None:
|
|
||||||
return found
|
|
||||||
return None
|
|
||||||
|
|
||||||
async def clients_covering(
|
async def clients_covering(
|
||||||
self, sources: list[str] | None = None
|
self, sources: list[str] | None = None
|
||||||
|
|
|
||||||
|
|
@ -27,23 +27,26 @@ def reported_database(client: "HaikuRAG", db_path: "Path | None") -> "Path | Non
|
||||||
return db_path if db_path is not None else client.store.db_path
|
return db_path if db_path is not None else client.store.db_path
|
||||||
|
|
||||||
|
|
||||||
async def database_lines(client: "HaikuRAG", db_path: Path) -> list[str]:
|
async def database_lines(client: "HaikuRAG") -> list[str]:
|
||||||
"""What one database reports about itself, without naming its location.
|
"""What one database reports about itself, without naming its location.
|
||||||
|
|
||||||
|
Reported through the connection the client already holds: a second one to the
|
||||||
|
same database would be a second open for the same statistics.
|
||||||
|
|
||||||
A failure is reported as a line rather than raised, so one unreachable
|
A failure is reported as a line rather than raised, so one unreachable
|
||||||
database does not cost the report on the others.
|
database does not cost the report on the others.
|
||||||
"""
|
"""
|
||||||
from haiku.rag.store.engine import ConnectionMode, connect_lancedb
|
from haiku.rag.store.engine import ConnectionMode
|
||||||
from haiku.rag.store.info import get_database_stats
|
from haiku.rag.store.info import get_database_stats
|
||||||
|
|
||||||
lines: list[str] = []
|
lines: list[str] = []
|
||||||
config = client.store._config
|
db_path = client.store.db_path
|
||||||
|
|
||||||
if client.store._connection_mode == ConnectionMode.LOCAL and not db_path.exists():
|
if client.store._connection_mode == ConnectionMode.LOCAL and not db_path.exists():
|
||||||
return ["[red]Database path does not exist.[/red]"]
|
return ["[red]Database path does not exist.[/red]"]
|
||||||
|
|
||||||
try:
|
try:
|
||||||
db = await connect_lancedb(config, db_path)
|
db = client.store.db
|
||||||
stats = await get_database_stats(db)
|
stats = await get_database_stats(db)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
return [f"[red]Failed to open database: {e}[/red]"]
|
return [f"[red]Failed to open database: {e}[/red]"]
|
||||||
|
|
@ -197,7 +200,7 @@ class InfoModal(ModalScreen):
|
||||||
lines.extend(block)
|
lines.extend(block)
|
||||||
else:
|
else:
|
||||||
lines.append(f"[bold $accent]path[/bold $accent]: {db_path}")
|
lines.append(f"[bold $accent]path[/bold $accent]: {db_path}")
|
||||||
lines.extend(await database_lines(self.client, db_path))
|
lines.extend(await database_lines(self.client))
|
||||||
|
|
||||||
lines.append("[bold]Versions[/bold]")
|
lines.append("[bold]Versions[/bold]")
|
||||||
versions = get_package_versions()
|
versions = get_package_versions()
|
||||||
|
|
@ -223,7 +226,7 @@ class InfoModal(ModalScreen):
|
||||||
# The client names a configured database by name and never by
|
# The client names a configured database by name and never by
|
||||||
# location, so its message is safe to show.
|
# location, so its message is safe to show.
|
||||||
return [*lines, f"[red]{e}[/red]", ""]
|
return [*lines, f"[red]{e}[/red]", ""]
|
||||||
return [*lines, *await database_lines(owner, owner.store.db_path)]
|
return [*lines, *await database_lines(owner)]
|
||||||
|
|
||||||
async def action_dismiss(self, result=None) -> None:
|
async def action_dismiss(self, result=None) -> None:
|
||||||
self.app.pop_screen()
|
self.app.pop_screen()
|
||||||
|
|
|
||||||
|
|
@ -369,6 +369,46 @@ class TestReportedDatabase:
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestReportingReusesTheConnection:
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_statistics_come_from_the_open_connection(
|
||||||
|
self, tmp_path, monkeypatch
|
||||||
|
):
|
||||||
|
"""The client already holds a connection to the database being reported,
|
||||||
|
so opening a second one would be an open for the same statistics."""
|
||||||
|
from haiku.rag.inspector.widgets.info_modal import database_lines
|
||||||
|
from haiku.rag.store.engine import ConnectionMode
|
||||||
|
|
||||||
|
def explode(*args, **kwargs):
|
||||||
|
raise AssertionError("opened a second connection to report statistics")
|
||||||
|
|
||||||
|
monkeypatch.setattr("haiku.rag.store.engine.connect_lancedb", explode)
|
||||||
|
|
||||||
|
asked: list[object] = []
|
||||||
|
|
||||||
|
async def fake_stats(db):
|
||||||
|
asked.append(db)
|
||||||
|
return {
|
||||||
|
"settings": {"exists": False},
|
||||||
|
"documents": {"num_rows": 1},
|
||||||
|
"document_meta": {"num_rows": 1},
|
||||||
|
"chunks": {"num_rows": 1},
|
||||||
|
}
|
||||||
|
|
||||||
|
monkeypatch.setattr("haiku.rag.store.info.get_database_stats", fake_stats)
|
||||||
|
|
||||||
|
connection = object()
|
||||||
|
client = MagicMock()
|
||||||
|
client.store.db = connection
|
||||||
|
client.store.db_path = tmp_path
|
||||||
|
client.store._connection_mode = ConnectionMode.LOCAL
|
||||||
|
|
||||||
|
lines = await database_lines(client)
|
||||||
|
|
||||||
|
assert asked == [connection]
|
||||||
|
assert any("documents" in line for line in lines)
|
||||||
|
|
||||||
|
|
||||||
class TestReportingEachDatabase:
|
class TestReportingEachDatabase:
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_a_database_that_cannot_be_opened_reports_itself(self):
|
async def test_a_database_that_cannot_be_opened_reports_itself(self):
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue