refactor(vector): define store contract
This commit is contained in:
@@ -0,0 +1,97 @@
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from tht.adapters.vector.thoth_http import ThothHttpVectorStore
|
||||
from tht.ports.vector import (
|
||||
VectorHit,
|
||||
VectorRecord,
|
||||
VectorStore,
|
||||
VectorWriteUnavailable,
|
||||
)
|
||||
|
||||
|
||||
def test_http_store_reports_reader_without_writer():
|
||||
reader = MagicMock()
|
||||
store = ThothHttpVectorStore(reader=reader, writer=None)
|
||||
|
||||
assert store.capabilities.search is True
|
||||
assert store.capabilities.upsert is False
|
||||
with pytest.raises(VectorWriteUnavailable):
|
||||
store.upsert("memory", [])
|
||||
|
||||
|
||||
def test_http_store_keeps_reader_and_writer_operations_separate():
|
||||
reader = MagicMock()
|
||||
reader.search_similar.return_value = [
|
||||
{
|
||||
"similarity": 0.75,
|
||||
"metadata": {
|
||||
"record_key": "m1",
|
||||
"kind": "memory",
|
||||
"ref": "session:s1",
|
||||
"title": "Choice",
|
||||
"content": "Use the curated table",
|
||||
},
|
||||
}
|
||||
]
|
||||
writer = MagicMock()
|
||||
writer.existing_hashes.return_value = {"m1": "abc"}
|
||||
writer.upsert_records.return_value = 1
|
||||
store = ThothHttpVectorStore(reader=reader, writer=writer)
|
||||
|
||||
hits = store.search(["memory"], [0.1, 0.2], limit=3, kinds=["memory"])
|
||||
assert hits == [
|
||||
VectorHit(
|
||||
id="m1",
|
||||
kind="memory",
|
||||
ref="session:s1",
|
||||
title="Choice",
|
||||
content="Use the curated table",
|
||||
metadata={
|
||||
"record_key": "m1",
|
||||
"kind": "memory",
|
||||
"ref": "session:s1",
|
||||
"title": "Choice",
|
||||
"content": "Use the curated table",
|
||||
},
|
||||
similarity=0.75,
|
||||
)
|
||||
]
|
||||
reader.search_similar.assert_called_once_with(
|
||||
"memory", [0.1, 0.2], 3, kinds=["memory"]
|
||||
)
|
||||
writer.search_similar.assert_not_called()
|
||||
|
||||
assert store.existing_hashes("memory", ["memory"]) == {"m1": "abc"}
|
||||
writer.existing_hashes.assert_called_once_with("memory", ["memory"])
|
||||
|
||||
records = [
|
||||
VectorRecord(
|
||||
id="m1",
|
||||
kind="memory",
|
||||
ref="session:s1",
|
||||
title="Choice",
|
||||
content="Use the curated table",
|
||||
metadata={"embedding": [0.1, 0.2], "content_hash": "abc"},
|
||||
)
|
||||
]
|
||||
assert store.upsert("memory", records) == 1
|
||||
writer.upsert_records.assert_called_once()
|
||||
reader.upsert_records.assert_not_called()
|
||||
|
||||
|
||||
def test_http_store_is_runtime_vector_store():
|
||||
store = ThothHttpVectorStore(reader=MagicMock(), writer=None)
|
||||
assert isinstance(store, VectorStore)
|
||||
|
||||
|
||||
def test_http_health_uses_reader_list_tables_and_reports_failure():
|
||||
reader = MagicMock()
|
||||
store = ThothHttpVectorStore(reader=reader, writer=None)
|
||||
assert store.health().ok is True
|
||||
|
||||
reader.list_tables.side_effect = RuntimeError("offline")
|
||||
health = store.health()
|
||||
assert health.ok is False
|
||||
assert health.detail == "offline"
|
||||
@@ -0,0 +1,6 @@
|
||||
"""Vector-store adapter implementations."""
|
||||
|
||||
from tht.adapters.vector.legacy_direct import LegacyDirectVectorStore
|
||||
from tht.adapters.vector.thoth_http import ThothHttpVectorStore
|
||||
|
||||
__all__ = ["LegacyDirectVectorStore", "ThothHttpVectorStore"]
|
||||
@@ -0,0 +1,47 @@
|
||||
"""Compatibility adapter for the existing direct PostgreSQL vector reader."""
|
||||
|
||||
from sqlalchemy import Engine
|
||||
|
||||
from tht.ports.vector import VectorCapabilities, VectorHealth, VectorRecord, VectorWriteUnavailable
|
||||
from tht.vectorstore.store import VectorHit, VectorStore as TableVectorStore
|
||||
|
||||
|
||||
class LegacyDirectVectorStore:
|
||||
"""Read-only port wrapper around the legacy table-scoped pgvector store."""
|
||||
|
||||
capabilities = VectorCapabilities(search=True, existing_hashes=False, upsert=False)
|
||||
|
||||
def __init__(self, engine: Engine, schema: str = "vectors", dim: int = 768):
|
||||
self._engine = engine
|
||||
self._schema = schema
|
||||
self._dim = dim
|
||||
|
||||
def health(self) -> VectorHealth:
|
||||
try:
|
||||
with self._engine.connect() as connection:
|
||||
connection.exec_driver_sql("SELECT 1")
|
||||
except Exception as exc:
|
||||
return VectorHealth(ok=False, detail=str(exc))
|
||||
return VectorHealth(ok=True)
|
||||
|
||||
def search(
|
||||
self,
|
||||
collections: list[str],
|
||||
embedding: list[float],
|
||||
*,
|
||||
limit: int,
|
||||
kinds: list[str] | None = None,
|
||||
) -> list[VectorHit]:
|
||||
hits: list[VectorHit] = []
|
||||
for collection in collections:
|
||||
table = TableVectorStore(
|
||||
self._engine, schema=self._schema, table=collection, dim=self._dim
|
||||
)
|
||||
hits.extend(table.search(embedding, top_n=limit, kinds=kinds))
|
||||
return sorted(hits, key=lambda hit: hit.similarity, reverse=True)[:limit]
|
||||
|
||||
def existing_hashes(self, collection: str, kinds: list[str]) -> dict[str, str]:
|
||||
raise VectorWriteUnavailable("Legacy direct reader has no writer interface")
|
||||
|
||||
def upsert(self, collection: str, records: list[VectorRecord]) -> int:
|
||||
raise VectorWriteUnavailable("Legacy direct reader has no writer interface")
|
||||
@@ -0,0 +1,93 @@
|
||||
"""Thoth vector HTTP adapter using distinct read and write clients."""
|
||||
|
||||
from tht.ports.vector import (
|
||||
VectorCapabilities,
|
||||
VectorHealth,
|
||||
VectorHit,
|
||||
VectorRecord,
|
||||
VectorStoreError,
|
||||
VectorWriteUnavailable,
|
||||
)
|
||||
from tht.vectorstore.rest_client import VectorRestClient
|
||||
from tht.vectorstore.store import content_hash, hit_from_metadata
|
||||
|
||||
|
||||
def _merge(hits: list[VectorHit], limit: int) -> list[VectorHit]:
|
||||
return sorted(hits, key=lambda hit: hit.similarity, reverse=True)[:limit]
|
||||
|
||||
|
||||
class ThothHttpVectorStore:
|
||||
"""Vector port backed by the existing allowlisted REST RPCs."""
|
||||
|
||||
def __init__(self, reader: VectorRestClient, writer: VectorRestClient | None):
|
||||
self._reader = reader
|
||||
self._writer = writer
|
||||
|
||||
@property
|
||||
def capabilities(self) -> VectorCapabilities:
|
||||
writable = self._writer is not None
|
||||
return VectorCapabilities(search=True, existing_hashes=writable, upsert=writable)
|
||||
|
||||
def health(self) -> VectorHealth:
|
||||
try:
|
||||
self._reader.list_tables()
|
||||
except Exception as exc:
|
||||
return VectorHealth(ok=False, detail=str(exc))
|
||||
return VectorHealth(ok=True)
|
||||
|
||||
def search(
|
||||
self,
|
||||
collections: list[str],
|
||||
embedding: list[float],
|
||||
*,
|
||||
limit: int,
|
||||
kinds: list[str] | None = None,
|
||||
) -> list[VectorHit]:
|
||||
hits: list[VectorHit] = []
|
||||
for collection in collections:
|
||||
rows = self._reader.search_similar(collection, embedding, limit, kinds=kinds)
|
||||
hits.extend(
|
||||
hit_from_metadata(row.get("similarity", 0.0), row.get("metadata"))
|
||||
for row in rows
|
||||
)
|
||||
if kinds:
|
||||
allowed = set(kinds)
|
||||
hits = [hit for hit in hits if hit.kind in allowed]
|
||||
return _merge(hits, limit)
|
||||
|
||||
def _require_writer(self) -> VectorRestClient:
|
||||
if self._writer is None:
|
||||
raise VectorWriteUnavailable("Vector writer credential is not configured")
|
||||
return self._writer
|
||||
|
||||
def existing_hashes(self, collection: str, kinds: list[str]) -> dict[str, str]:
|
||||
return self._require_writer().existing_hashes(collection, kinds)
|
||||
|
||||
def upsert(self, collection: str, records: list[VectorRecord]) -> int:
|
||||
writer = self._require_writer()
|
||||
rows = [self._row(record) for record in records]
|
||||
return writer.upsert_records(collection, rows)
|
||||
|
||||
@staticmethod
|
||||
def _row(record: VectorRecord) -> dict:
|
||||
extra = dict(record.metadata)
|
||||
try:
|
||||
embedding = extra.pop("embedding")
|
||||
except KeyError as exc:
|
||||
raise VectorStoreError(f"Vector record {record.id!r} has no embedding") from exc
|
||||
digest = extra.pop("content_hash", content_hash(record.content))
|
||||
metadata = {
|
||||
"kind": record.kind,
|
||||
"ref": record.ref,
|
||||
"record_key": record.id,
|
||||
"title": record.title,
|
||||
"content": record.content,
|
||||
**extra,
|
||||
}
|
||||
return {
|
||||
"record_key": record.id,
|
||||
"kind": record.kind,
|
||||
"content_hash": digest,
|
||||
"metadata": metadata,
|
||||
"embedding": embedding,
|
||||
}
|
||||
@@ -0,0 +1,60 @@
|
||||
"""Transport-neutral vector-store contract and canonical vector models."""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Protocol, runtime_checkable
|
||||
|
||||
from tht.vectorstore.records import VectorRecord
|
||||
from tht.vectorstore.store import VectorHit
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class VectorCapabilities:
|
||||
search: bool = True
|
||||
existing_hashes: bool = False
|
||||
upsert: bool = False
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class VectorHealth:
|
||||
ok: bool
|
||||
detail: str | None = None
|
||||
|
||||
|
||||
class VectorStoreError(Exception):
|
||||
"""Base error exposed by vector adapters."""
|
||||
|
||||
|
||||
class VectorWriteUnavailable(VectorStoreError):
|
||||
"""Raised when a deployment has no vector writer credential."""
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class VectorStore(Protocol):
|
||||
@property
|
||||
def capabilities(self) -> VectorCapabilities: ...
|
||||
|
||||
def health(self) -> VectorHealth: ...
|
||||
|
||||
def search(
|
||||
self,
|
||||
collections: list[str],
|
||||
embedding: list[float],
|
||||
*,
|
||||
limit: int,
|
||||
kinds: list[str] | None = None,
|
||||
) -> list[VectorHit]: ...
|
||||
|
||||
def existing_hashes(self, collection: str, kinds: list[str]) -> dict[str, str]: ...
|
||||
|
||||
def upsert(self, collection: str, records: list[VectorRecord]) -> int: ...
|
||||
|
||||
|
||||
__all__ = [
|
||||
"VectorCapabilities",
|
||||
"VectorHealth",
|
||||
"VectorHit",
|
||||
"VectorRecord",
|
||||
"VectorStore",
|
||||
"VectorStoreError",
|
||||
"VectorWriteUnavailable",
|
||||
]
|
||||
@@ -9,8 +9,10 @@ Entrambe mappano i `kind` sulle tabelle per-dominio dello schema `vectors`.
|
||||
|
||||
from sqlalchemy import Engine
|
||||
|
||||
from tht.adapters.vector.legacy_direct import LegacyDirectVectorStore
|
||||
from tht.adapters.vector.thoth_http import ThothHttpVectorStore
|
||||
from tht.vectorstore.rest_client import VectorRestClient
|
||||
from tht.vectorstore.store import VectorHit, VectorStore, hit_from_metadata
|
||||
from tht.vectorstore.store import VectorHit
|
||||
|
||||
# kind Thoth → tabella dello schema `vectors`.
|
||||
KIND_TO_TABLE = {
|
||||
@@ -30,30 +32,19 @@ def tables_for_kinds(kinds: list[str] | None) -> list[str]:
|
||||
return sorted({KIND_TO_TABLE[k] for k in kinds if k in KIND_TO_TABLE})
|
||||
|
||||
|
||||
def _merge(hits: list[VectorHit], top_n: int) -> list[VectorHit]:
|
||||
return sorted(hits, key=lambda h: h.similarity, reverse=True)[:top_n]
|
||||
|
||||
|
||||
class RestSearcher:
|
||||
"""Similarity search via REST: una chiamata `search_similar` per tabella, poi fusione."""
|
||||
|
||||
def __init__(self, client: VectorRestClient):
|
||||
self.client = client
|
||||
self._store = ThothHttpVectorStore(reader=client, writer=None)
|
||||
|
||||
def search(
|
||||
self, query_vec: list[float], top_n: int = 10, kinds: list[str] | None = None
|
||||
) -> list[VectorHit]:
|
||||
hits: list[VectorHit] = []
|
||||
for table in tables_for_kinds(kinds):
|
||||
for row in self.client.search_similar(table, query_vec, top_n, kinds=kinds):
|
||||
hits.append(hit_from_metadata(row.get("similarity", 0.0), row.get("metadata")))
|
||||
# Il filtro per kind avviene server-side (RPC con `kinds`); il post-filter resta
|
||||
# come difesa per il fallback legacy (server pre-migrazione: 404 -> query senza
|
||||
# filtro) e per parita' col path diretto (#25).
|
||||
if kinds:
|
||||
allowed = set(kinds)
|
||||
hits = [h for h in hits if h.kind in allowed]
|
||||
return _merge(hits, top_n)
|
||||
return self._store.search(
|
||||
tables_for_kinds(kinds), query_vec, limit=top_n, kinds=kinds
|
||||
)
|
||||
|
||||
|
||||
class DirectSearcher:
|
||||
@@ -63,13 +54,11 @@ class DirectSearcher:
|
||||
self.engine = engine
|
||||
self.schema = schema
|
||||
self.dim = dim
|
||||
self._store = LegacyDirectVectorStore(engine, schema=schema, dim=dim)
|
||||
|
||||
def search(
|
||||
self, query_vec: list[float], top_n: int = 10, kinds: list[str] | None = None
|
||||
) -> list[VectorHit]:
|
||||
hits: list[VectorHit] = []
|
||||
for table in tables_for_kinds(kinds):
|
||||
store = VectorStore(self.engine, schema=self.schema, table=table, dim=self.dim)
|
||||
# passa kinds: dentro schema_records filtra schema_table vs schema_column (#25).
|
||||
hits.extend(store.search(query_vec, top_n=top_n, kinds=kinds))
|
||||
return _merge(hits, top_n)
|
||||
return self._store.search(
|
||||
tables_for_kinds(kinds), query_vec, limit=top_n, kinds=kinds
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user