"""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 import re 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.connect_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, metadata_filter: dict | None = None, ) -> list[dict]: """Ricerca per similarità coseno su `vectors.`: 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 metadata_filter is not None: # ACTIVE corpus reads must never degrade to an unfiltered legacy RPC: # filtering after LIMIT is incomplete and could expose stale generations. return self._call( "search_similar", {**args, "kinds": kinds, "metadata_filter": metadata_filter}, ) or [] 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) def delete_generation(self, table_name: str, generation: str, workspace_id: str) -> int: if table_name != "evidence" or re.fullmatch(r"gen:[0-9a-f]{32}", generation) is None: raise ValueError("generation must be canonical") if re.fullmatch(r"[a-z][a-z0-9_-]{0,63}", workspace_id) is None: raise ValueError("workspace namespace must be canonical") try: payload = self._call( "delete_vector_generation", {"table_name": table_name, "kind": "evidence", "generation": generation, "workspace_id": workspace_id}, ) except VectorRestError as error: if "HTTP 404" in str(error): raise VectorRestError( "delete_vector_generation RPC is unavailable; deploy the cleanup migration" ) from None raise if isinstance(payload, dict): return int(payload.get("deleted", 0)) return 0 def list_evidence_generations(self, table_name: str, workspace_id: str) -> list[str]: if re.fullmatch(r"[a-z][a-z0-9_-]{0,63}", workspace_id) is None: raise ValueError("workspace namespace must be canonical") try: rows = self._call( "list_evidence_generations", {"table_name": table_name, "kind": "evidence", "workspace_id": workspace_id}, ) or [] except VectorRestError as error: if "HTTP 404" in str(error): raise VectorRestError( "list_evidence_generations RPC is unavailable; deploy the cleanup migration" ) from None raise if not isinstance(rows, list) or any( not isinstance(row, dict) or re.fullmatch(r"gen:[0-9a-f]{32}", str(row.get("generation", ""))) is None for row in rows ): raise VectorRestError("list_evidence_generations returned malformed data") return sorted({row["generation"] for row in rows})