192 lines
7.1 KiB
Python
192 lines
7.1 KiB
Python
import hashlib
|
|
import json
|
|
from dataclasses import dataclass
|
|
|
|
from sqlalchemy import Engine, text
|
|
|
|
from tht.vectorstore.records import VectorRecord
|
|
|
|
|
|
def content_hash(content: str) -> str:
|
|
return hashlib.sha256(content.encode()).hexdigest()
|
|
|
|
|
|
def _to_vector_literal(vec: list[float]) -> str:
|
|
return "[" + ",".join(f"{x:.8f}" for x in vec) + "]"
|
|
|
|
|
|
@dataclass
|
|
class SyncStats:
|
|
added: int = 0
|
|
updated: int = 0
|
|
deleted: int = 0
|
|
unchanged: int = 0
|
|
|
|
|
|
@dataclass
|
|
class VectorHit:
|
|
id: str
|
|
kind: str
|
|
ref: str
|
|
title: str
|
|
content: str
|
|
metadata: dict
|
|
similarity: float
|
|
|
|
|
|
def hit_from_metadata(similarity: float, metadata: dict | None) -> VectorHit:
|
|
"""Ricostruisce un VectorHit dal solo `metadata` (più la similarity). È l'unico modo
|
|
disponibile leggendo via REST (`search_similar` ritorna id/similarity/metadata), e viene
|
|
usato anche dalla lettura diretta per avere un'unica logica. Tollerante: usa default sui
|
|
campi assenti (es. metadata estranei della tabella fake remota)."""
|
|
md = metadata or {}
|
|
return VectorHit(
|
|
id=md.get("record_key", ""),
|
|
kind=md.get("record_kind", md.get("kind", "")),
|
|
ref=md.get("ref", ""),
|
|
title=md.get("title", ""),
|
|
content=md.get("content", ""),
|
|
metadata=md,
|
|
similarity=float(similarity),
|
|
)
|
|
|
|
|
|
class VectorStore:
|
|
"""Tabella pgvector table-scoped: scrittura diretta (loading) su una tabella dello schema
|
|
`vectors`. La lettura via REST avviene su `search_similar`; questo store serve al loading e
|
|
alla lettura diretta (dev/test). Il contratto della tabella remota richiede `id` (BIGSERIAL),
|
|
`embedding vector(N)` e `metadata jsonb`; le colonne extra (`record_key`, `kind`,
|
|
`content_hash`) servono solo al loader e non sono esposte dalla REST."""
|
|
|
|
def __init__(
|
|
self, engine: Engine, schema: str = "vectors", table: str = "records", dim: int = 768
|
|
):
|
|
self.engine = engine
|
|
self.schema = schema
|
|
self.dim = dim
|
|
self._table = f"{schema}.{table}"
|
|
|
|
def init_schema(self) -> None:
|
|
with self.engine.begin() as conn:
|
|
conn.execute(text("CREATE EXTENSION IF NOT EXISTS vector"))
|
|
conn.execute(text(f"CREATE SCHEMA IF NOT EXISTS {self.schema}"))
|
|
conn.execute(text(f"""
|
|
CREATE TABLE IF NOT EXISTS {self._table} (
|
|
id bigserial PRIMARY KEY,
|
|
record_key text UNIQUE NOT NULL,
|
|
kind text NOT NULL,
|
|
content_hash text NOT NULL,
|
|
metadata jsonb NOT NULL DEFAULT '{{}}',
|
|
embedding vector({self.dim}) NOT NULL,
|
|
indexed_at timestamptz NOT NULL DEFAULT now()
|
|
)
|
|
"""))
|
|
conn.execute(text(
|
|
f"CREATE INDEX IF NOT EXISTS {self._idx('embedding')} ON {self._table} "
|
|
f"USING hnsw (embedding vector_cosine_ops) WITH (m = 16, ef_construction = 200)"
|
|
))
|
|
conn.execute(text(
|
|
f"CREATE INDEX IF NOT EXISTS {self._idx('kind')} ON {self._table} (kind)"
|
|
))
|
|
# GRANT al ruolo di sola lettura della REST, solo se esiste (assente in test/locale).
|
|
conn.execute(text(f"""
|
|
DO $$ BEGIN
|
|
IF EXISTS (SELECT 1 FROM pg_roles WHERE rolname = 'vector_reader') THEN
|
|
EXECUTE 'GRANT SELECT ON {self._table} TO vector_reader';
|
|
END IF;
|
|
END $$;
|
|
"""))
|
|
|
|
def _idx(self, suffix: str) -> str:
|
|
return f"{self._table.replace('.', '_')}_{suffix}_idx"
|
|
|
|
def clear(self) -> None:
|
|
with self.engine.begin() as conn:
|
|
conn.execute(text(f"DELETE FROM {self._table}"))
|
|
|
|
def existing_hashes(self, kinds: set[str]) -> dict[str, str]:
|
|
q = text(
|
|
f"SELECT record_key, content_hash FROM {self._table} WHERE kind = ANY(:kinds)"
|
|
)
|
|
with self.engine.connect() as conn:
|
|
return dict(conn.execute(q, {"kinds": list(kinds)}).fetchall())
|
|
|
|
def sync(self, records: list[VectorRecord], embedder, kinds: set[str]) -> SyncStats:
|
|
"""Allinea l'indice ai record correnti (per i kind dati): embedda solo il nuovo
|
|
o il modificato, elimina cio' che non esiste piu'."""
|
|
stats = SyncStats()
|
|
existing = self.existing_hashes(kinds)
|
|
current_ids = {r.id for r in records}
|
|
|
|
to_embed: list[VectorRecord] = []
|
|
for r in records:
|
|
h = content_hash(r.content)
|
|
if r.id not in existing:
|
|
to_embed.append(r)
|
|
stats.added += 1
|
|
elif existing[r.id] != h:
|
|
to_embed.append(r)
|
|
stats.updated += 1
|
|
else:
|
|
stats.unchanged += 1
|
|
|
|
vectors = embedder.embed_documents([r.content for r in to_embed]) if to_embed else []
|
|
|
|
upsert = text(f"""
|
|
INSERT INTO {self._table}
|
|
(record_key, kind, content_hash, metadata, embedding)
|
|
VALUES
|
|
(:record_key, :kind, :content_hash, CAST(:metadata AS jsonb),
|
|
CAST(:embedding AS vector))
|
|
ON CONFLICT (record_key) DO UPDATE SET
|
|
kind = EXCLUDED.kind, content_hash = EXCLUDED.content_hash,
|
|
metadata = EXCLUDED.metadata, embedding = EXCLUDED.embedding,
|
|
indexed_at = now()
|
|
""")
|
|
stale = [i for i in existing if i not in current_ids]
|
|
with self.engine.begin() as conn:
|
|
for r, vec in zip(to_embed, vectors):
|
|
conn.execute(upsert, {
|
|
"record_key": r.id, "kind": r.kind,
|
|
"content_hash": content_hash(r.content),
|
|
"metadata": json.dumps(_pack_metadata(r)),
|
|
"embedding": _to_vector_literal(vec),
|
|
})
|
|
if stale:
|
|
conn.execute(
|
|
text(f"DELETE FROM {self._table} WHERE record_key = ANY(:ids)"),
|
|
{"ids": stale},
|
|
)
|
|
stats.deleted = len(stale)
|
|
return stats
|
|
|
|
def search(
|
|
self, query_vec: list[float], top_n: int = 10, kinds: list[str] | None = None
|
|
) -> list[VectorHit]:
|
|
where = "WHERE kind = ANY(:kinds)" if kinds else ""
|
|
q = text(f"""
|
|
SELECT metadata, 1 - (embedding <=> CAST(:q AS vector)) AS similarity
|
|
FROM {self._table}
|
|
{where}
|
|
ORDER BY embedding <=> CAST(:q AS vector)
|
|
LIMIT :top_n
|
|
""")
|
|
params: dict = {"q": _to_vector_literal(query_vec), "top_n": top_n}
|
|
if kinds:
|
|
params["kinds"] = kinds
|
|
with self.engine.connect() as conn:
|
|
rows = conn.execute(q, params).fetchall()
|
|
return [hit_from_metadata(r.similarity, r.metadata) for r in rows]
|
|
|
|
|
|
def _pack_metadata(r: VectorRecord) -> dict:
|
|
"""Impacchetta nel `metadata` (unica colonna letta via REST) tutta la semantica Thoth."""
|
|
return {
|
|
"kind": r.kind,
|
|
"ref": r.ref,
|
|
"record_key": r.id,
|
|
"title": r.title,
|
|
"content": r.content,
|
|
**r.metadata,
|
|
}
|