| """ |
| 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 |
|
|
| |
| |
| |
| 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 |
|
|
| |
| _response_cache: Dict[str, Tuple[str, float]] = {} |
| _CACHE_TTL_SECONDS = 3600 |
|
|
|
|
| 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 |
| |
| 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()) |
| |
| 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 = "", |
| ): |
| """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" |
|
|
| |
| _qdrant_client = client |
| _collection_name = collection_name |
|
|
| |
| 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 |
|
|
| |
| 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 |
|
|
| |
| 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 |
|
|
| |
| 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) |
| |
| 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: |
| |
| 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 = [] |
| _reranker = None |
|
|
|
|
| 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()), |
| ) |
| ] |
| ) |
| |
| |
| 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 [] |
|
|
|
|