RICS / app /vectorstore /qdrant_wrapper.py
StormShadow308's picture
Speed up generation and ingestion; add live report progress on /status.
c893230
Raw
History Blame Contribute Delete
18.3 kB
"""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)