Files
ThothII/harness/tht/adapters/vector/thoth_http.py
T

88 lines
2.9 KiB
Python

"""Thoth vector HTTP adapter using distinct read and write clients."""
from tht.ports.vector import (
VectorCapabilities,
VectorHealth,
VectorHit,
VectorWriteRecord,
VectorWriteUnavailable,
)
from tht.vectorstore.rest_client import VectorRestClient
from tht.vectorstore.store import 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[VectorWriteRecord]) -> int:
writer = self._require_writer()
rows = [self._row(record) for record in records]
return writer.upsert_records(collection, rows)
@staticmethod
def _row(write_record: VectorWriteRecord) -> dict:
record = write_record.record
metadata = {
"kind": record.kind,
"ref": record.ref,
"record_key": record.id,
"title": record.title,
"content": record.content,
**record.metadata,
}
return {
"record_key": record.id,
"kind": record.kind,
"content_hash": write_record.content_hash,
"metadata": metadata,
"embedding": write_record.embedding,
}