121 lines
4.8 KiB
Python
121 lines
4.8 KiB
Python
"""Client per la similarity search del pgvector esposta via Supabase/PostgREST.
|
|
|
|
Endpoint dedicato (es. https://host/vector/v1/), distinto dal DWH. La lettura usa
|
|
`search_similar`; la scrittura remota usa RPC allowlist con una API key separata.
|
|
Errori in italiano e azionabili, stile `rest/client.py`.
|
|
"""
|
|
|
|
import requests
|
|
|
|
from tht.config import RestConfig
|
|
|
|
|
|
class VectorRestError(Exception):
|
|
"""Errore di accesso al vector store via REST, con messaggio leggibile per il reviewer."""
|
|
|
|
|
|
class VectorRestClient:
|
|
def __init__(self, cfg: RestConfig):
|
|
self.cfg = cfg
|
|
self._base = cfg.base_url.rstrip("/")
|
|
|
|
@property
|
|
def api_key(self) -> str:
|
|
"""The REST API key for this client (spec D11: reader and writer carry
|
|
distinct keys against the same endpoint)."""
|
|
return self.cfg.api_key
|
|
|
|
def _post(self, fn: str, args: dict) -> requests.Response:
|
|
url = f"{self._base}/rpc/{fn}"
|
|
verify: bool | str = self.cfg.ssl_ca if self.cfg.ssl_ca else True
|
|
try:
|
|
return requests.post(
|
|
url,
|
|
json=args,
|
|
headers={"X-API-Key": self.cfg.api_key},
|
|
timeout=self.cfg.timeout,
|
|
verify=verify,
|
|
)
|
|
except requests.RequestException as e:
|
|
raise VectorRestError(
|
|
f"Vector REST non raggiungibile su {self.cfg.base_url} (rpc {fn}): {e}"
|
|
) from e
|
|
|
|
def _error_msg(self, fn: str, resp: requests.Response) -> str:
|
|
try:
|
|
body = resp.json()
|
|
detail = body.get("message") or body.get("details") or resp.text
|
|
except Exception:
|
|
detail = resp.text
|
|
return f"Vector REST rpc {fn} → HTTP {resp.status_code}: {detail}"
|
|
|
|
def _call(self, fn: str, args: dict):
|
|
resp = self._post(fn, args)
|
|
if not resp.ok:
|
|
raise VectorRestError(self._error_msg(fn, resp))
|
|
if resp.status_code == 204 or not resp.text:
|
|
return None
|
|
return resp.json()
|
|
|
|
def search_similar(
|
|
self, table_name: str, query_embedding: list[float], limit_count: int,
|
|
kinds: list[str] | None = None,
|
|
) -> list[dict]:
|
|
"""Ricerca per similarità coseno su `vectors.<table_name>`: ritorna le righe
|
|
`{id, similarity, metadata}` ordinate per similarity decrescente. Con `kinds`
|
|
il filtro avviene server-side nel WHERE della RPC (evita la diluizione del
|
|
top-k quando piu' kind condividono la tabella, es. memory/solved_question).
|
|
Su un server legacy senza il parametro (PostgREST 404) ritenta senza filtro:
|
|
resta il post-filter client-side di RestSearcher."""
|
|
args = {
|
|
"query_embedding": query_embedding,
|
|
"limit_count": limit_count,
|
|
"table_name": table_name,
|
|
}
|
|
if kinds is not None:
|
|
try:
|
|
return self._call("search_similar", {**args, "kinds": kinds}) or []
|
|
except VectorRestError as e:
|
|
if "HTTP 404" not in str(e):
|
|
raise
|
|
# funzione a 3 argomenti (pre-migrazione kinds): fallback senza filtro
|
|
return self._call("search_similar", args) or []
|
|
|
|
def list_tables(self) -> list[dict]:
|
|
"""Tabelle vettoriali disponibili: `{table_name, vector_dimensions, …}`."""
|
|
return self._call("list_tables", {}) or []
|
|
|
|
def existing_hashes(self, table_name: str, kinds: list[str]) -> dict[str, str]:
|
|
"""Hash correnti per sync incrementale su una tabella vector allowlisted.
|
|
|
|
RPC attesa: `existing_vector_hashes(table_name, kinds)` -> righe
|
|
`{record_key, content_hash}`.
|
|
"""
|
|
rows = self._call(
|
|
"existing_vector_hashes",
|
|
{"table_name": table_name, "kinds": kinds},
|
|
) or []
|
|
return {row["record_key"]: row["content_hash"] for row in rows}
|
|
|
|
def upsert_records(self, table_name: str, rows: list[dict]) -> int:
|
|
"""Upsert controllato di record vettoriali già embeddati.
|
|
|
|
RPC attesa: `upsert_vector_records(table_name, rows)` -> `{upserted: N}` o righe.
|
|
Non espone delete/clear: il cleanup distruttivo resta solo-server.
|
|
"""
|
|
payload = self._call(
|
|
"upsert_vector_records",
|
|
{"table_name": table_name, "rows": rows},
|
|
)
|
|
if payload is None:
|
|
return len(rows)
|
|
if isinstance(payload, dict):
|
|
return int(payload.get("upserted", len(rows)))
|
|
# PostgREST puo' incapsulare uno scalar jsonb in una lista [{"upserted": N}]:
|
|
# estrai il conteggio dal primo elemento invece di restituire len(lista)=1.
|
|
if isinstance(payload, list):
|
|
if payload and isinstance(payload[0], dict) and "upserted" in payload[0]:
|
|
return int(payload[0]["upserted"])
|
|
return len(payload)
|
|
return len(rows)
|