""" DeepMed-AI — tools/vector_store.py Qdrant Cloud vector store: embeddings, creation, FILTERED retrieval, and response cache. Key features: - search_by_drug_name(): Qdrant payload filtering for 100% drug retrieval - get_retriever(): Fast mode (BM25 + Vector, k=15) - get_deep_retriever(): Deep mode (BM25 + Vector k=25 → CrossEncoder Reranker top 5) """ import hashlib import os import time from typing import Dict, List, Optional, Tuple # ── Fix HuggingFace cache path trên HF Spaces ───────────────────────────────── # HF Spaces đặt HOME=/nonexistent → sentence-transformers crash khi tải model. # Buộc tất cả cache về /tmp (luôn writable trong mọi container). os.environ.setdefault("HF_HOME", "/tmp/huggingface") os.environ.setdefault("HUGGINGFACE_HUB_CACHE", "/tmp/huggingface/hub") os.environ.setdefault("SENTENCE_TRANSFORMERS_HOME", "/tmp/sentence-transformers") os.environ.setdefault("TRANSFORMERS_CACHE", "/tmp/transformers") os.environ.setdefault("XDG_CACHE_HOME", "/tmp/.cache") from langchain_core.documents import Document from app.core.logging_config import logger _embeddings = None _vectorstore = None _qdrant_client = None _collection_name = None # ── Simple in-memory response cache with TTL ─────────────────────────────────── _response_cache: Dict[str, Tuple[str, float]] = {} _CACHE_TTL_SECONDS = 3600 # 1 hour def get_cached_response(question: str) -> Optional[str]: """Return cached response for question if available and not expired.""" key = hashlib.md5(question.lower().strip().encode()).hexdigest() if key in _response_cache: response, ts = _response_cache[key] if time.time() - ts < _CACHE_TTL_SECONDS: logger.info("Cache HIT for question (key: %s...)", key[:8]) return response # Expired — remove del _response_cache[key] return None def set_cached_response(question: str, response: str) -> None: """Cache a response for a given question.""" key = hashlib.md5(question.lower().strip().encode()).hexdigest() _response_cache[key] = (response, time.time()) # Simple eviction: keep max 200 entries if len(_response_cache) > 200: oldest_key = min(_response_cache, key=lambda k: _response_cache[k][1]) del _response_cache[oldest_key] def get_embeddings(): """Return a cached HuggingFace sentence-transformer embeddings instance.""" global _embeddings if _embeddings is None: from langchain_huggingface.embeddings import HuggingFaceEmbeddings _embeddings = HuggingFaceEmbeddings( model_name="sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2" ) logger.info("Embeddings model loaded (paraphrase-multilingual-MiniLM-L12-v2)") return _embeddings def get_or_create_vectorstore( documents: Optional[List[Document]] = None, persist_dir: str = "", # Unused for cloud, kept for compatibility ): """Load Qdrant Vector DB or create new collection and push documents.""" global _vectorstore, _qdrant_client, _collection_name if _vectorstore is not None: return _vectorstore from qdrant_client import QdrantClient from qdrant_client.models import Distance, VectorParams from app.core.config import QDRANT_API_KEY, QDRANT_URL embeddings = get_embeddings() if not QDRANT_URL or not QDRANT_API_KEY: logger.warning( "QDRANT_URL or QDRANT_API_KEY not configured. RAG will be disabled." ) return None logger.info("Connecting to Qdrant Cloud Cluster at %s...", QDRANT_URL) client = QdrantClient( url=QDRANT_URL, api_key=QDRANT_API_KEY, timeout=30, ) collection_name = "deepmed_rag_v5" # v5: metadata filtering for 100% drug retrieval # Save for later use by search_by_drug_name _qdrant_client = client _collection_name = collection_name # Check if collection exists try: collections = client.get_collections().collections collection_exists = any(c.name == collection_name for c in collections) except Exception as e: logger.error("Failed to connect Qdrant: %s", e) return None # Tạo collection NẾU chưa tồn tại (kể cả không có documents) if not collection_exists: logger.info("Creating Qdrant collection: %s", collection_name) try: client.create_collection( collection_name=collection_name, vectors_config=VectorParams(size=384, distance=Distance.COSINE), ) logger.info("Collection '%s' created successfully.", collection_name) except Exception as e: logger.error("Failed to create Qdrant collection: %s", e) return None # Khởi tạo VectorStore — bắt lỗi để không crash 500 try: from langchain_qdrant import QdrantVectorStore _vectorstore = QdrantVectorStore( client=client, collection_name=collection_name, embedding=embeddings, ) except ImportError: logger.warning("langchain-qdrant not found, falling back to langchain_community.Qdrant") try: from langchain_community.vectorstores import Qdrant _vectorstore = Qdrant( client=client, collection_name=collection_name, embeddings=embeddings, ) except Exception as e: logger.error("Failed to init vectorstore fallback: %s", e) return None except Exception as e: logger.error("Failed to init QdrantVectorStore: %s", e) return None # ── Upload documents ──────────────────────────────────────────────────── def _batch_upload(docs_to_upload): """Upload documents in small batches to avoid Qdrant Cloud timeouts.""" BATCH_SIZE = 100 total = len(docs_to_upload) for i in range(0, total, BATCH_SIZE): batch = docs_to_upload[i:i + BATCH_SIZE] try: _vectorstore.add_documents(batch) logger.info("Uploaded batch %d/%d (%d docs)", i // BATCH_SIZE + 1, (total + BATCH_SIZE - 1) // BATCH_SIZE, len(batch)) except Exception as e: logger.error("Failed to upload batch %d: %s", i // BATCH_SIZE + 1, e) # Continue with next batch — don't lose everything logger.info("Upload complete! Total: %d docs", total) if documents and not collection_exists: logger.info("Uploading %d documents to Qdrant Cloud (batched)...", len(documents)) _batch_upload(documents) elif documents and collection_exists: # Check if collection is empty before skipping try: count_result = client.count(collection_name=collection_name) if count_result.count == 0: logger.info("Collection exists but EMPTY — uploading %d docs...", len(documents)) _batch_upload(documents) else: logger.info("Qdrant collection '%s' has %d vectors. Skip re-upload.", collection_name, count_result.count) except Exception as e: logger.warning("Could not check collection count: %s. Skipping upload.", e) return _vectorstore _all_splits = [] # Keep document splits in memory for BM25 _reranker = None # CrossEncoder Reranker singleton def _get_reranker(top_n: int = 5): """Return a cached CrossEncoder Reranker (BGE-reranker-v2-m3). Used in Deep mode: re-scores candidate docs for clinical precision. Lazy-loaded on first call to avoid startup overhead if only Fast mode is used. """ global _reranker if _reranker is None: try: from langchain_community.cross_encoders import HuggingFaceCrossEncoder from langchain.retrievers.document_compressors import CrossEncoderReranker logger.info("Loading Reranker Model (BGE-reranker-v2-m3)...") reranker_model = HuggingFaceCrossEncoder( model_name="BAAI/bge-reranker-v2-m3", ) _reranker = CrossEncoderReranker(model=reranker_model, top_n=top_n) logger.info("Reranker loaded: BGE-reranker-v2-m3, top_n=%d", top_n) except Exception as e: logger.error("Failed to load Reranker: %s", e) return None return _reranker def _build_ensemble(k: int): """Build a BM25 + Vector EnsembleRetriever with the given k.""" vs = get_or_create_vectorstore() if not vs: return None vector_retriever = vs.as_retriever(search_kwargs={"k": k}) if _all_splits and len(_all_splits) > 10: try: from langchain_community.retrievers import BM25Retriever from langchain.retrievers.ensemble import EnsembleRetriever bm25_retriever = BM25Retriever.from_documents(_all_splits) bm25_retriever.k = k ensemble = EnsembleRetriever( retrievers=[bm25_retriever, vector_retriever], weights=[0.5, 0.5], ) logger.info("Hybrid ensemble: BM25(%d docs) + Vector, k=%d", len(_all_splits), k) return ensemble except Exception as e: logger.warning("BM25 init failed, falling back to vector-only: %s", e) logger.info("Vector-only retriever, k=%d", k) return vector_retriever def get_retriever(k: int = 15): """Return a FAST HYBRID retriever: BM25 keyword + Qdrant vector, 50/50 ensemble. Fast mode (default): - BM25: exact keyword matching (catches drug names, hoạt chất perfectly) - Vector: semantic similarity (catches paraphrased/conceptual questions) - Combined via EnsembleRetriever with equal weights - k=15: returns 15 candidate docs — good balance of speed & coverage """ return _build_ensemble(k) def get_deep_retriever(k: int = 25, top_n: int = 5): """Return a DEEP retriever: BM25 + Vector (k=25) → CrossEncoder Reranker (top_n=5). Deep mode: - Casts a wide net with k=25 candidate docs - CrossEncoder (BGE-reranker-v2-m3) re-scores each candidate against the query - Returns only the top_n=5 most relevant docs — clinical-grade precision - Slower but significantly more accurate for complex protocol queries """ base_retriever = _build_ensemble(k) if not base_retriever: return None reranker = _get_reranker(top_n=top_n) if not reranker: logger.warning("Reranker unavailable, falling back to fast retriever") return base_retriever try: from langchain.retrievers import ContextualCompressionRetriever deep = ContextualCompressionRetriever( base_compressor=reranker, base_retriever=base_retriever, ) logger.info("Deep retriever: Ensemble(k=%d) → Reranker(top_n=%d)", k, top_n) return deep except Exception as e: logger.error("Failed to build deep retriever: %s", e) return base_retriever def store_splits_for_bm25(splits: list): """Store document splits in memory for BM25 retriever. Called from main.py after document loading. """ global _all_splits _all_splits = splits logger.info("BM25: stored %d splits in memory", len(splits)) def search_by_drug_name(drug_name: str, query: str, k: int = 8) -> List[Document]: """Search for chunks of a SPECIFIC drug using Qdrant metadata filter. This is the KEY function that guarantees 100% retrieval accuracy. Instead of relying on embedding similarity (which may return wrong drug docs), this filters by metadata.drug_name first, then ranks by semantic similarity. Args: drug_name: The drug name to filter by (e.g., "MIDANTIN") query: The user's question (for semantic ranking within filtered results) k: Number of results to return Returns: List of Documents from ONLY the specified drug's .md file """ global _vectorstore, _qdrant_client, _collection_name if not _vectorstore or not _qdrant_client: logger.warning("search_by_drug_name: vectorstore not ready") return [] try: from qdrant_client.models import Filter, FieldCondition, MatchValue drug_filter = Filter( must=[ FieldCondition( key="metadata.drug_name", match=MatchValue(value=drug_name.upper()), ) ] ) # Use the vectorstore's similarity_search with filter results = _vectorstore.similarity_search( query=query, k=k, filter=drug_filter, ) if results: logger.info("DrugFilter: Found %d chunks for drug '%s'", len(results), drug_name) else: logger.info("DrugFilter: No filtered results for '%s', will fallback to semantic", drug_name) return results except Exception as e: logger.error("DrugFilter search failed for '%s': %s", drug_name, e) return [] def search_by_doc_type(doc_type: str, query: str, k: int = 5) -> List[Document]: """Search within a specific doc type (drug_info, reference_pdf, etc.).""" global _vectorstore if not _vectorstore: return [] try: from qdrant_client.models import Filter, FieldCondition, MatchValue type_filter = Filter( must=[ FieldCondition( key="metadata.doc_type", match=MatchValue(value=doc_type), ) ] ) return _vectorstore.similarity_search(query=query, k=k, filter=type_filter) except Exception as e: logger.error("DocType search failed: %s", e) return [] def search_by_ingredient(ingredient_keyword: str, query: str, k: int = 8) -> List[Document]: """Search for drugs containing a specific ACTIVE INGREDIENT. Uses Qdrant filter on metadata.ingredient_keywords (list field). When a Qdrant field is a list, MatchValue matches if ANY element equals the value. Example: ingredient_keyword="AMOXICILIN" will find all drugs containing amoxicilin, such as MIDANTIN (Amoxicilin+acid clavulanic), FABAMOX (Amoxicillin), etc. Args: ingredient_keyword: Uppercase ingredient name (e.g., "AMOXICILIN", "CEFTRIAXON") query: The user's question (for semantic ranking within filtered results) k: Number of results to return """ global _vectorstore if not _vectorstore: logger.warning("search_by_ingredient: vectorstore not ready") return [] try: from qdrant_client.models import Filter, FieldCondition, MatchValue ingredient_filter = Filter( must=[ FieldCondition( key="metadata.ingredient_keywords", match=MatchValue(value=ingredient_keyword.upper()), ) ] ) results = _vectorstore.similarity_search( query=query, k=k, filter=ingredient_filter, ) if results: logger.info("IngredientFilter: Found %d chunks for ingredient '%s'", len(results), ingredient_keyword) else: logger.info("IngredientFilter: No results for ingredient '%s'", ingredient_keyword) return results except Exception as e: logger.error("IngredientFilter search failed for '%s': %s", ingredient_keyword, e) return []