DeepMedAI / backend /app /tools /vector_store.py
PBThuong's picture
Thiết lập lại thư viện y khoa sạch và cập nhật chroma_db
8eaa451
Raw
History Blame Contribute Delete
16.1 kB
"""
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 []