"""Embedding denso (Ollama/llama.cpp) e sparse BM25, con catena di fallback resiliente. Catena da EMBED_CHAIN (JSON, stesso formato di RERANK_CHAIN + campo "api"): il primo nodo raggiungibile vince, i nodi falliti entrano in cooldown. Se EMBED_CHAIN è vuota si usa il comportamento legacy (endpoint singolo da EMBED_API/EMBED_URL/EMBED_API_KEY). L'ultimo errore viene rilanciato al client come gli endpoint precedenti: nessuna degradazione silenziosa della scrittura. """ from __future__ import annotations import json import time from dataclasses import dataclass from typing import Any, Optional import httpx from qdrant_client.http import models as qm import metrics from config import ( EMBED_API, EMBED_API_KEY, EMBED_CHAIN, EMBED_DIM, EMBED_MODEL, EMBED_RETRY_COOLDOWN_S, EMBED_TIMEOUT_MS, EMBED_URL, SPARSE_VECTOR_NAME, log, ) try: from fastembed import SparseTextEmbedding _sparse_model: Optional[Any] = None SPARSE_AVAILABLE = True except Exception: # noqa: BLE001 _sparse_model = None SPARSE_AVAILABLE = False log.warning("fastembed non disponibile: hybrid retrieval disattivato") _http: Optional[httpx.AsyncClient] = None @dataclass(frozen=True) class EmbedNode: """Un endpoint embedding nella catena di fallback.""" name: str url: str api: str # "llamacpp" (/v1/embeddings) | "ollama" (/api/embed) key: str timeout_ms: int def parse_chain(raw: str, default_api: str, default_url: str, default_key: str) -> list[EmbedNode]: """Parsa EMBED_CHAIN (JSON); vuota o invalida → endpoint legacy singolo.""" nodes: list[EmbedNode] = [] if raw: try: entries = json.loads(raw) for entry in entries if isinstance(entries, list) else []: if not isinstance(entry, dict) or not entry.get("url"): continue url = str(entry["url"]).rstrip("/") api = str(entry.get("api") or "llamacpp") if api not in ("llamacpp", "ollama") or not url.startswith(("http://", "https://")): continue nodes.append( EmbedNode( name=str(entry.get("name") or url), url=url, api=api, key=str(entry.get("key") or ""), timeout_ms=int(entry.get("timeout_ms", EMBED_TIMEOUT_MS)), ) ) except (json.JSONDecodeError, TypeError, ValueError): log.error("EMBED_CHAIN non è JSON valido: uso l'endpoint legacy") if not nodes and default_url: # Compatibilità legacy: endpoint singolo dagli env EMBED_* nodes = [EmbedNode(name="embed", url=default_url.rstrip("/"), api=default_api, key=default_key, timeout_ms=EMBED_TIMEOUT_MS)] return nodes _chain: Optional[list[EmbedNode]] = None _down_until: dict[str, float] = {} # url → monotonic deadline del cooldown _http: Optional[httpx.AsyncClient] = None def _get_chain() -> list[EmbedNode]: global _chain if _chain is None: _chain = parse_chain(EMBED_CHAIN, EMBED_API, EMBED_URL, EMBED_API_KEY) return _chain def reset_chain_cache() -> None: """Forza il re-parse della catena (usato dai test).""" global _chain _chain = None _down_until.clear() def get_http() -> httpx.AsyncClient: global _http if _http is None: _http = httpx.AsyncClient(timeout=30) return _http async def close_http() -> None: global _http if _http is not None: await _http.aclose() _http = None def chain_nodes() -> list[EmbedNode]: return _get_chain() async def embed(text: str) -> list[float]: """Embedding con catena di fallback: ritorna il vettore o rilancia dopo l'ultimo fallimento.""" chain = _get_chain() if not chain: raise RuntimeError("nessun endpoint embedding configurato") now = time.monotonic() live = [n for n in chain if _down_until.get(n.url, 0) <= now] if not live: # tutti in cooldown: ritenta comunque il primo (meglio di un fallimento immediato) live = [chain[0]] payload = {"model": EMBED_MODEL, "input": text} started = time.monotonic() last_exc: Optional[Exception] = None for node in live: headers = {"Content-Type": "application/json"} if node.key: headers["Authorization"] = f"Bearer {node.key}" path = "/v1/embeddings" if node.api == "llamacpp" else "/api/embed" try: t0 = time.monotonic() response = await get_http().post( f"{node.url}{path}", json=payload, headers=headers, timeout=httpx.Timeout(node.timeout_ms / 1000.0), ) response.raise_for_status() data = response.json() vector = data["data"][0]["embedding"] if node.api == "llamacpp" else data["embeddings"][0] if len(vector) != EMBED_DIM: raise ValueError(f"dimensione vettore {len(vector)} != EMBED_DIM {EMBED_DIM}") took = int((time.monotonic() - started) * 1000) metrics.record_embed(node.name, True, took) return vector except (httpx.HTTPError, ValueError, KeyError, IndexError, TypeError) as exc: took = int((time.monotonic() - t0) * 1000) _down_until[node.url] = time.monotonic() + EMBED_RETRY_COOLDOWN_S last_exc = exc metrics.record_embed(node.name, False, took) log.warning( "embed: nodo '%s' fallito dopo %dms (%s: %s) → cooldown %ds", node.name, took, exc.__class__.__name__, exc, EMBED_RETRY_COOLDOWN_S, ) raise RuntimeError(f"tutti i nodi embedding falliti ({len(live)} tentativi)") from last_exc def get_sparse_model(): global _sparse_model if _sparse_model is None and SPARSE_AVAILABLE: _sparse_model = SparseTextEmbedding(model_name="Qdrant/bm25") return _sparse_model def sparse_encode(text: str) -> Optional[qm.SparseVector]: model = get_sparse_model() if model is None: return None emb = next(model.embed(text)) return qm.SparseVector(indices=emb.indices.tolist(), values=emb.values.tolist()) def backfill_sparse(qdrant: Any, collection: str) -> None: if not SPARSE_AVAILABLE: return offset: Any = None updated = 0 while True: points, next_offset = qdrant.scroll( collection_name=collection, limit=100, with_payload=["text"], with_vectors=True, offset=offset, ) batch: list[qm.PointStruct] = [] for p in points: vecs = p.vector or {} if SPARSE_VECTOR_NAME in vecs: continue text = (p.payload or {}).get("text", "") if not text: continue sparse = sparse_encode(text) if sparse is None: continue batch.append(qm.PointStruct(id=p.id, vector={SPARSE_VECTOR_NAME: sparse})) if batch: # update_vectors: aggiorna SOLO il vettore sparso, preservando payload e vettore denso # (upsert parziale sostituirebbe l'intero punto — incidente 2026-08-16) qdrant.update_vectors(collection_name=collection, points=batch) updated += len(batch) if not next_offset: break offset = next_offset if updated: log.info("backfill sparse: %d record aggiornati", updated)