"""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"), document_purpose=payload.get("document_purpose") or "report_source", 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 _post_filter_purpose( self, rows: list[SearchResult], *, purpose_in: frozenset[str] | None, exclude_purpose: frozenset[str] | None, ) -> list[SearchResult]: """Apply document_purpose filters; KB rows are always exempt.""" if purpose_in is None and exclude_purpose is None: return rows out: list[SearchResult] = [] for r in rows: if r.kb: out.append(r) continue purpose = r.document_purpose or "report_source" if purpose_in is not None and purpose not in purpose_in: continue if exclude_purpose is not None and purpose in exclude_purpose: continue 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, purpose_in: frozenset[str] | None = None, exclude_purpose: frozenset[str] | None = None, ) -> list[SearchResult]: use_hybrid = self._want_hybrid(tenant_id) # Over-fetch so the purpose filter has enough candidates to keep `k` # results after dropping the style_corpus rows (or, in the inverse # case, after dropping report_source rows from the style-only call). fetch_k = k if (purpose_in is None and exclude_purpose is None) else max(k * 3, k + 20) if use_hybrid: rows = self._hybrid_search( query, tenant_id, fetch_k, hierarchy_level=hierarchy_level, doc_id_in=doc_id_in, ) else: rows = self._vector_search( query, tenant_id, fetch_k, hierarchy_level=hierarchy_level, doc_id_in=doc_id_in, ) rows = self._post_filter_purpose( rows, purpose_in=purpose_in, exclude_purpose=exclude_purpose ) return rows[:k] 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, purpose_in: frozenset[str] | None = None, exclude_purpose: frozenset[str] | None = None, ) -> list[SearchResult]: vector = await run_sync_in_executor(self._embedding.embed_query, query) # Over-fetch when a purpose filter is in play so we still get `k` # rows after dropping the wrong-purpose candidates. purpose_fanout = purpose_in is not None or exclude_purpose is not None fetch_k = max(k * (6 if purpose_fanout else 4), k + (30 if purpose_fanout else 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] rows = self._post_filter_purpose( rows, purpose_in=purpose_in, exclude_purpose=exclude_purpose ) 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)