"""Semantic cache for retrieval results (Qdrant + query embeddings).""" from __future__ import annotations import json import logging import time import uuid from typing import TYPE_CHECKING from app.async_executor import run_sync_in_executor from app.config import settings from app.models.schemas import SearchResult if TYPE_CHECKING: from qdrant_client import AsyncQdrantClient from qdrant_client.http import models as qmodels logger = logging.getLogger(__name__) _cache_instance: "SemanticCache | None" = None def reset_semantic_cache() -> None: """Reset semantic cache singleton (tests).""" global _cache_instance _cache_instance = None def is_semantic_cache_active() -> bool: """True when Qdrant-backed semantic cache is enabled.""" return _use_semantic_cache() def _use_semantic_cache() -> bool: return bool( settings.semantic_cache_enabled and (settings.vectorstore_backend or "faiss").strip().lower() == "qdrant" ) class SemanticCache: """Cache retrieval results keyed by tenant + query embedding similarity.""" def __init__( self, client: "AsyncQdrantClient", collection: str, *, embed_query, ttl_seconds: int, threshold: float, ) -> None: self._client = client self._collection = collection self._embed_query = embed_query self._ttl_seconds = ttl_seconds self._threshold = threshold self._ready = False async def _embed_async(self, query: str) -> list[float]: """Embed without blocking the event loop (sentence-transformers / OpenAI).""" return await run_sync_in_executor(self._embed_query, query) async def ensure_collection(self, vector_size: int) -> None: if self._ready: return from qdrant_client.http import models as qmodels names = {c.name for c in (await self._client.get_collections()).collections} if self._collection not in names: await self._client.create_collection( collection_name=self._collection, vectors_config=qmodels.VectorParams( size=vector_size, distance=qmodels.Distance.COSINE, ), hnsw_config=qmodels.HnswConfigDiff(m=16, ef_construct=100), on_disk_payload=True, ) self._ready = True async def get(self, query: str, tenant_id: str) -> list[SearchResult] | None: from qdrant_client.http import models as qmodels vector = await self._embed_async(query) await self.ensure_collection(len(vector)) cutoff = time.time() - self._ttl_seconds hits = await self._client.search( collection_name=self._collection, query_vector=vector, limit=3, query_filter=qmodels.Filter( must=[ qmodels.FieldCondition( key="tenant_id", match=qmodels.MatchValue(value=tenant_id), ), ] ), ) for hit in hits: if float(hit.score) < self._threshold: continue created = float((hit.payload or {}).get("created_at", 0)) if created < cutoff: continue raw = (hit.payload or {}).get("results_json", "[]") try: rows = json.loads(raw) return [SearchResult.model_validate(r) for r in rows] except Exception as exc: # noqa: BLE001 logger.warning("Semantic cache payload decode failed: %s", exc) return None async def put( self, query: str, tenant_id: str, results: list[SearchResult], ) -> None: from qdrant_client.http import models as qmodels if not results: return vector = await self._embed_async(query) await self.ensure_collection(len(vector)) payload = { "tenant_id": tenant_id, "query": query[:500], "created_at": time.time(), "results_json": json.dumps([r.model_dump() for r in results]), } await self._client.upsert( collection_name=self._collection, points=[ qmodels.PointStruct( id=str(uuid.uuid4()), vector=vector, payload=payload, ) ], ) async def invalidate_tenant(self, tenant_id: str) -> None: """Drop cached retrieval rows for a tenant (after ingest/delete).""" from qdrant_client.http import models as qmodels if not tenant_id: return vector = await self._embed_async("cache dimension probe") await self.ensure_collection(len(vector)) await self._client.delete( collection_name=self._collection, points_selector=qmodels.FilterSelector( filter=qmodels.Filter( must=[ qmodels.FieldCondition( key="tenant_id", match=qmodels.MatchValue(value=tenant_id), ) ] ) ), ) logger.debug("Semantic cache invalidated for tenant=%s", tenant_id) async def invalidate_semantic_cache_for_tenant(tenant_id: str) -> None: """Best-effort semantic cache invalidation (no-op when cache disabled).""" cache = get_semantic_cache() if cache is None: return try: await cache.invalidate_tenant(tenant_id) except Exception as exc: # noqa: BLE001 logger.warning("Semantic cache invalidation failed tenant=%s: %s", tenant_id, exc) def get_semantic_cache() -> SemanticCache | None: """Return a semantic cache singleton when enabled; otherwise ``None``.""" global _cache_instance if not _use_semantic_cache(): return None if _cache_instance is not None: return _cache_instance from app.embeddings.factory import get_embedding_client from app.vectorstore.qdrant_async import get_async_qdrant_client embedding = get_embedding_client() client = get_async_qdrant_client() _cache_instance = SemanticCache( client=client, collection=settings.qdrant_cache_collection, embed_query=embedding.embed_query, ttl_seconds=int(settings.semantic_cache_ttl_hours) * 3600, threshold=float(settings.semantic_cache_similarity_threshold), ) return _cache_instance