Spaces:
Sleeping
Sleeping
| """Qdrant vector store with tenant filters, optional hybrid BM25+RRF, and async search.""" | |
| from __future__ import annotations | |
| import logging | |
| import threading | |
| import uuid | |
| from typing import Any | |
| from langchain_core.documents import Document | |
| from langchain_core.embeddings import Embeddings | |
| from qdrant_client import QdrantClient | |
| from qdrant_client.http import models as qmodels | |
| from app.async_executor import run_sync_in_executor | |
| from app.config import settings | |
| from app.models.schemas import SearchResult | |
| from app.retrieval.hybrid import bm25_search, reciprocal_rank_fusion | |
| from app.vectorstore.base import VectorStore | |
| logger = logging.getLogger(__name__) | |
| _RICS_CHUNK_NS = uuid.uuid5(uuid.NAMESPACE_DNS, "rics-uk-project/chunk") | |
| def _stable_point_id(chunk_id: str) -> str: | |
| """Deterministic Qdrant point id (UUID) from arbitrary chunk_id strings.""" | |
| return str(uuid.uuid5(_RICS_CHUNK_NS, chunk_id)) | |
| def _is_qdrant_backend() -> bool: | |
| return (settings.vectorstore_backend or "faiss").strip().lower() == "qdrant" | |
| def _doc_to_payload(doc: Document) -> dict[str, Any]: | |
| meta = dict(doc.metadata or {}) | |
| return { | |
| "chunk_id": str(meta.get("chunk_id", "")), | |
| "doc_id": str(meta.get("doc_id", "")), | |
| "tenant_id": str(meta.get("tenant_id", "")), | |
| "hierarchy_level": str(meta.get("hierarchy_level") or "paragraph"), | |
| "section_type": str(meta.get("section_type", "general")), | |
| "section_title": meta.get("section_title"), | |
| "section_id": meta.get("section_id"), | |
| "paragraph_index": meta.get("paragraph_index"), | |
| "parent_chunk_id": meta.get("parent_chunk_id"), | |
| "source": meta.get("source"), | |
| "kb": meta.get("kb"), | |
| "kb_path": meta.get("kb_path"), | |
| "chunk_role": meta.get("chunk_role"), | |
| "text": doc.page_content, | |
| } | |
| def _payload_to_result(payload: dict[str, Any], score: float) -> SearchResult: | |
| return SearchResult( | |
| chunk_id=str(payload.get("chunk_id", "")), | |
| doc_id=str(payload.get("doc_id", "")), | |
| tenant_id=str(payload.get("tenant_id", "")), | |
| text=str(payload.get("text", "")), | |
| score=float(score), | |
| section_type=str(payload.get("section_type", "general")), | |
| hierarchy_level=str(payload.get("hierarchy_level") or "paragraph"), | |
| section_title=payload.get("section_title"), | |
| section_id=payload.get("section_id"), | |
| paragraph_index=payload.get("paragraph_index"), | |
| parent_chunk_id=payload.get("parent_chunk_id"), | |
| source=payload.get("source"), | |
| kb=payload.get("kb"), | |
| kb_path=payload.get("kb_path"), | |
| chunk_role=payload.get("chunk_role"), | |
| ) | |
| class QdrantVectorStore(VectorStore): | |
| """Tenant-scoped Qdrant store with HNSW and optional hybrid retrieval.""" | |
| def __init__(self, embedding: Embeddings) -> None: | |
| self._embedding = embedding | |
| self._client = QdrantClient( | |
| url=settings.qdrant_url, | |
| api_key=settings.qdrant_api_key or None, | |
| ) | |
| self._collection = settings.qdrant_collection | |
| self._lock = threading.RLock() | |
| self._bm25_texts: dict[str, list[str]] = {} | |
| self._bm25_rows: dict[str, list[SearchResult]] = {} | |
| self._vector_size: int | None = None | |
| self._bm25_loaded: set[str] = set() | |
| self._ensure_collection() | |
| def _get_async_client(self) -> Any: | |
| from app.vectorstore.qdrant_async import get_async_qdrant_client | |
| return get_async_qdrant_client() | |
| def _embed_dim(self) -> int: | |
| if self._vector_size is None: | |
| vec = self._embedding.embed_query("dimension probe") | |
| self._vector_size = len(vec) | |
| return self._vector_size | |
| def _ensure_collection(self) -> None: | |
| dim = self._embed_dim() | |
| names = {c.name for c in self._client.get_collections().collections} | |
| if self._collection in names: | |
| return | |
| self._client.create_collection( | |
| collection_name=self._collection, | |
| vectors_config=qmodels.VectorParams(size=dim, distance=qmodels.Distance.COSINE), | |
| hnsw_config=qmodels.HnswConfigDiff(m=16, ef_construct=100), | |
| on_disk_payload=True, | |
| ) | |
| self._client.create_payload_index( | |
| collection_name=self._collection, | |
| field_name="tenant_id", | |
| field_schema=qmodels.PayloadSchemaType.KEYWORD, | |
| ) | |
| logger.info("Created Qdrant collection %s (dim=%d)", self._collection, dim) | |
| def _tenant_filter( | |
| self, | |
| tenant_id: str, | |
| *, | |
| hierarchy_level: str | None, | |
| doc_id_in: frozenset[str] | None, | |
| ) -> qmodels.Filter: | |
| must: list[qmodels.Condition] = [ | |
| qmodels.FieldCondition( | |
| key="tenant_id", | |
| match=qmodels.MatchValue(value=tenant_id), | |
| ) | |
| ] | |
| if hierarchy_level is not None: | |
| must.append( | |
| qmodels.FieldCondition( | |
| key="hierarchy_level", | |
| match=qmodels.MatchValue(value=hierarchy_level), | |
| ) | |
| ) | |
| return qmodels.Filter(must=must) | |
| def _post_filter_doc_ids( | |
| self, | |
| rows: list[SearchResult], | |
| doc_id_in: frozenset[str] | None, | |
| ) -> list[SearchResult]: | |
| if doc_id_in is None: | |
| return rows | |
| out: list[SearchResult] = [] | |
| for r in rows: | |
| if r.doc_id in doc_id_in or r.kb: | |
| out.append(r) | |
| return out | |
| def _vector_search( | |
| self, | |
| query: str, | |
| tenant_id: str, | |
| k: int, | |
| *, | |
| hierarchy_level: str | None, | |
| doc_id_in: frozenset[str] | None, | |
| ) -> list[SearchResult]: | |
| vector = self._embedding.embed_query(query) | |
| fetch_k = max(k * 4, k + 10) | |
| hits = self._client.search( | |
| collection_name=self._collection, | |
| query_vector=vector, | |
| limit=fetch_k, | |
| query_filter=self._tenant_filter( | |
| tenant_id, | |
| hierarchy_level=hierarchy_level, | |
| doc_id_in=None, | |
| ), | |
| ) | |
| rows = [ | |
| _payload_to_result(hit.payload or {}, float(hit.score)) | |
| for hit in hits | |
| if (hit.payload or {}).get("tenant_id") == tenant_id | |
| ] | |
| return self._post_filter_doc_ids(rows, doc_id_in)[:k] | |
| def _hybrid_search( | |
| self, | |
| query: str, | |
| tenant_id: str, | |
| k: int, | |
| *, | |
| hierarchy_level: str | None, | |
| doc_id_in: frozenset[str] | None, | |
| ) -> list[SearchResult]: | |
| vector_hits = self._vector_search( | |
| query, | |
| tenant_id, | |
| max(k * 3, 30), | |
| hierarchy_level=hierarchy_level, | |
| doc_id_in=doc_id_in, | |
| ) | |
| texts = self._bm25_texts.get(tenant_id, []) | |
| rows = self._bm25_rows.get(tenant_id, []) | |
| if hierarchy_level is not None: | |
| filtered = [ | |
| (t, r) | |
| for t, r in zip(texts, rows, strict=True) | |
| if r.hierarchy_level == hierarchy_level | |
| ] | |
| if filtered: | |
| texts, rows = [x[0] for x in filtered], [x[1] for x in filtered] | |
| if doc_id_in is not None: | |
| filtered = [ | |
| (t, r) | |
| for t, r in zip(texts, rows, strict=True) | |
| if r.doc_id in doc_id_in or r.kb | |
| ] | |
| if filtered: | |
| texts, rows = [x[0] for x in filtered], [x[1] for x in filtered] | |
| bm25_hits = bm25_search(query, texts=texts, meta_rows=rows, k=max(k * 3, 30)) | |
| return reciprocal_rank_fusion([vector_hits, bm25_hits], top_n=k) | |
| def _rebuild_bm25_for_tenant(self, tenant_id: str) -> bool: | |
| """Load BM25 corpus for ``tenant_id`` from Qdrant scroll (survives restarts).""" | |
| texts: list[str] = [] | |
| rows: list[SearchResult] = [] | |
| offset: Any = None | |
| while True: | |
| batch, offset = self._client.scroll( | |
| collection_name=self._collection, | |
| scroll_filter=self._tenant_filter(tenant_id, hierarchy_level=None, doc_id_in=None), | |
| limit=256, | |
| offset=offset, | |
| with_payload=True, | |
| with_vectors=False, | |
| ) | |
| for rec in batch: | |
| payload = rec.payload or {} | |
| texts.append(str(payload.get("text", ""))) | |
| rows.append(_payload_to_result(payload, 0.0)) | |
| if offset is None: | |
| break | |
| with self._lock: | |
| self._bm25_texts[tenant_id] = texts | |
| self._bm25_rows[tenant_id] = rows | |
| if texts: | |
| self._bm25_loaded.add(tenant_id) | |
| if texts: | |
| logger.debug("Rebuilt BM25 index for tenant=%s (%d chunks)", tenant_id, len(texts)) | |
| return bool(texts) | |
| def _ensure_bm25_for_tenant(self, tenant_id: str) -> bool: | |
| with self._lock: | |
| if self._bm25_rows.get(tenant_id): | |
| return True | |
| return self._rebuild_bm25_for_tenant(tenant_id) | |
| def _want_hybrid(self, tenant_id: str) -> bool: | |
| return bool( | |
| settings.enable_hybrid_retrieval | |
| and _is_qdrant_backend() | |
| and self._ensure_bm25_for_tenant(tenant_id) | |
| ) | |
| def _update_bm25(self, documents: list[Document]) -> None: | |
| for doc in documents: | |
| meta = doc.metadata or {} | |
| tenant = str(meta.get("tenant_id", "")) | |
| if not tenant: | |
| continue | |
| chunk_id = str(meta.get("chunk_id", "")) | |
| payload = _doc_to_payload(doc) | |
| row = _payload_to_result(payload, 0.0) | |
| texts = self._bm25_texts.setdefault(tenant, []) | |
| rows = self._bm25_rows.setdefault(tenant, []) | |
| if chunk_id: | |
| for i, existing in enumerate(rows): | |
| if existing.chunk_id == chunk_id: | |
| texts[i] = doc.page_content | |
| rows[i] = row | |
| break | |
| else: | |
| texts.append(doc.page_content) | |
| rows.append(row) | |
| else: | |
| texts.append(doc.page_content) | |
| rows.append(row) | |
| def add_documents(self, documents: list[Document]) -> None: | |
| if not documents: | |
| return | |
| points: list[qmodels.PointStruct] = [] | |
| for doc in documents: | |
| meta = doc.metadata or {} | |
| chunk_id = str(meta.get("chunk_id") or uuid.uuid4()) | |
| payload = _doc_to_payload(doc) | |
| vector = self._embedding.embed_documents([doc.page_content])[0] | |
| points.append( | |
| qmodels.PointStruct( | |
| id=_stable_point_id(chunk_id), | |
| vector=vector, | |
| payload=payload, | |
| ) | |
| ) | |
| with self._lock: | |
| self._client.upsert(collection_name=self._collection, points=points) | |
| self._update_bm25(documents) | |
| logger.debug("Upserted %d points into Qdrant", len(points)) | |
| def search( | |
| self, | |
| query: str, | |
| tenant_id: str, | |
| k: int = 10, | |
| *, | |
| hierarchy_level: str | None = None, | |
| doc_id_in: frozenset[str] | None = None, | |
| ) -> list[SearchResult]: | |
| use_hybrid = self._want_hybrid(tenant_id) | |
| if use_hybrid: | |
| return self._hybrid_search( | |
| query, | |
| tenant_id, | |
| k, | |
| hierarchy_level=hierarchy_level, | |
| doc_id_in=doc_id_in, | |
| ) | |
| return self._vector_search( | |
| query, | |
| tenant_id, | |
| k, | |
| hierarchy_level=hierarchy_level, | |
| doc_id_in=doc_id_in, | |
| ) | |
| def _hybrid_merge_vector_rows( | |
| self, | |
| query: str, | |
| tenant_id: str, | |
| k: int, | |
| rows: list[SearchResult], | |
| *, | |
| hierarchy_level: str | None, | |
| doc_id_in: frozenset[str] | None, | |
| ) -> list[SearchResult]: | |
| """BM25 + RRF merge for pre-fetched vector rows (CPU-bound; run off event loop).""" | |
| texts = self._bm25_texts.get(tenant_id, []) | |
| meta = self._bm25_rows.get(tenant_id, []) | |
| if hierarchy_level is not None: | |
| pairs = [ | |
| (t, r) | |
| for t, r in zip(texts, meta, strict=True) | |
| if r.hierarchy_level == hierarchy_level | |
| ] | |
| if pairs: | |
| texts, meta = [p[0] for p in pairs], [p[1] for p in pairs] | |
| if doc_id_in is not None: | |
| pairs = [ | |
| (t, r) | |
| for t, r in zip(texts, meta, strict=True) | |
| if r.doc_id in doc_id_in or r.kb | |
| ] | |
| if pairs: | |
| texts, meta = [p[0] for p in pairs], [p[1] for p in pairs] | |
| bm25_hits = bm25_search(query, texts=texts, meta_rows=meta, k=max(k * 3, 30)) | |
| return reciprocal_rank_fusion([rows, bm25_hits], top_n=k) | |
| async def search_async( | |
| self, | |
| query: str, | |
| tenant_id: str, | |
| k: int = 10, | |
| *, | |
| hierarchy_level: str | None = None, | |
| doc_id_in: frozenset[str] | None = None, | |
| ) -> list[SearchResult]: | |
| vector = await run_sync_in_executor(self._embedding.embed_query, query) | |
| fetch_k = max(k * 4, k + 10) | |
| client = self._get_async_client() | |
| hits = await client.search( | |
| collection_name=self._collection, | |
| query_vector=vector, | |
| limit=fetch_k, | |
| query_filter=self._tenant_filter( | |
| tenant_id, | |
| hierarchy_level=hierarchy_level, | |
| doc_id_in=None, | |
| ), | |
| ) | |
| rows = [ | |
| _payload_to_result(hit.payload or {}, float(hit.score)) | |
| for hit in hits | |
| if (hit.payload or {}).get("tenant_id") == tenant_id | |
| ] | |
| rows = self._post_filter_doc_ids(rows, doc_id_in)[:fetch_k] | |
| if self._want_hybrid(tenant_id): | |
| return await run_sync_in_executor( | |
| self._hybrid_merge_vector_rows, | |
| query, | |
| tenant_id, | |
| k, | |
| rows, | |
| hierarchy_level=hierarchy_level, | |
| doc_id_in=doc_id_in, | |
| ) | |
| return rows[:k] | |
| def delete_document(self, doc_id: str) -> None: | |
| with self._lock: | |
| self._client.delete( | |
| collection_name=self._collection, | |
| points_selector=qmodels.FilterSelector( | |
| filter=qmodels.Filter( | |
| must=[ | |
| qmodels.FieldCondition( | |
| key="doc_id", | |
| match=qmodels.MatchValue(value=doc_id), | |
| ) | |
| ] | |
| ) | |
| ), | |
| ) | |
| for tenant in list(self._bm25_rows.keys()): | |
| rows = self._bm25_rows[tenant] | |
| texts = self._bm25_texts[tenant] | |
| kept = [(t, r) for t, r in zip(texts, rows, strict=True) if r.doc_id != doc_id] | |
| self._bm25_texts[tenant] = [x[0] for x in kept] | |
| self._bm25_rows[tenant] = [x[1] for x in kept] | |
| if not self._bm25_rows[tenant]: | |
| self._bm25_loaded.discard(tenant) | |
| def count(self, tenant_id: str) -> int: | |
| result = self._client.count( | |
| collection_name=self._collection, | |
| count_filter=qmodels.Filter( | |
| must=[ | |
| qmodels.FieldCondition( | |
| key="tenant_id", | |
| match=qmodels.MatchValue(value=tenant_id), | |
| ) | |
| ] | |
| ), | |
| ) | |
| return int(result.count) | |
| def count_for_doc(self, doc_id: str) -> int: | |
| result = self._client.count( | |
| collection_name=self._collection, | |
| count_filter=qmodels.Filter( | |
| must=[ | |
| qmodels.FieldCondition( | |
| key="doc_id", | |
| match=qmodels.MatchValue(value=doc_id), | |
| ) | |
| ] | |
| ), | |
| ) | |
| return int(result.count) | |