import os import time import pickle from pathlib import Path from qdrant_client import QdrantClient from qdrant_client.models import Distance, VectorParams, PointStruct from rank_bm25 import BM25Okapi from tqdm import tqdm from underthesea import word_tokenize from src.embedder import embed_texts def _tokenize_vi(text: str) -> list[str]: """Tách từ tiếng Việt, fallback về split nếu lỗi.""" try: return word_tokenize(text.lower(), format="text").split() except Exception: return text.lower().split() QDRANT_PATH = "qdrant_storage" COLLECTION_NAME = "rag_docs" BM25_PATH = "bm25_index.pkl" BATCH_SIZE = 20 VECTOR_DIM = 3072 # gemini-embedding-001 def get_client(): url = os.getenv("QDRANT_URL") api_key = os.getenv("QDRANT_API_KEY") if url and api_key: return QdrantClient(url=url, api_key=api_key, timeout=60) return QdrantClient(path=QDRANT_PATH) def get_or_create_collection(client): existing = [c.name for c in client.get_collections().collections] if COLLECTION_NAME not in existing: client.create_collection( collection_name=COLLECTION_NAME, vectors_config=VectorParams(size=VECTOR_DIM, distance=Distance.COSINE), ) return client.get_collection(COLLECTION_NAME) CHECKPOINT_PATH = "embed_checkpoint.pkl" def add_chunks(chunks: list[dict]): """Embed và lưu chunks vào Qdrant + build BM25 index. Có checkpoint để resume khi lỗi quota.""" client = get_client() get_or_create_collection(client) texts = [c["text"] for c in chunks] metadatas = [c["metadata"] for c in chunks] # Load checkpoint nếu có done_batches: dict[int, list] = {} if Path(CHECKPOINT_PATH).exists(): with open(CHECKPOINT_PATH, "rb") as f: done_batches = pickle.load(f) print(f"Resume từ checkpoint: đã xong {len(done_batches)} batches (~{len(done_batches) * BATCH_SIZE} chunks)") pending = [i for i in range(0, len(chunks), BATCH_SIZE) if i not in done_batches] print(f"Embedding {len(chunks)} chunks... ({len(pending)} batches còn lại)") for i in tqdm(pending): batch_texts = texts[i:i + BATCH_SIZE] for attempt in range(3): try: embeddings = embed_texts(batch_texts) break except Exception as e: if "429" in str(e) or "RESOURCE_EXHAUSTED" in str(e): wait = 60 * (attempt + 1) print(f"\n[Quota] Lỗi rate limit, đợi {wait}s rồi thử lại (lần {attempt + 1}/3)...") time.sleep(wait) if attempt == 2: print("[Quota] Hết quota ngày hôm nay. Chạy lại vào ngày mai, sẽ resume từ checkpoint.") raise else: raise done_batches[i] = embeddings with open(CHECKPOINT_PATH, "wb") as f: pickle.dump(done_batches, f) time.sleep(BATCH_SIZE * 1.5) # Ghép tất cả embeddings theo thứ tự all_embeddings = [] for i in range(0, len(chunks), BATCH_SIZE): all_embeddings.extend(done_batches[i]) # Lưu vào Qdrant theo batch để tránh timeout UPSERT_BATCH = 200 print(f"Upserting {len(all_embeddings)} points vào Qdrant...") for i in tqdm(range(0, len(all_embeddings), UPSERT_BATCH)): batch_points = [ PointStruct(id=i + j, vector=all_embeddings[i + j], payload={"text": texts[i + j], **metadatas[i + j]}) for j in range(min(UPSERT_BATCH, len(all_embeddings) - i)) ] client.upsert(collection_name=COLLECTION_NAME, points=batch_points) # Build BM25 index print("Building BM25 index...") tokenized = [_tokenize_vi(t) for t in texts] bm25 = BM25Okapi(tokenized) with open(BM25_PATH, "wb") as f: pickle.dump({"bm25": bm25, "texts": texts, "metadatas": metadatas}, f) # Xóa checkpoint sau khi hoàn thành Path(CHECKPOINT_PATH).unlink(missing_ok=True) print(f"Đã lưu {len(chunks)} chunks vào Qdrant + BM25") def _build_bm25_from_qdrant() -> dict: """Rebuild BM25 index từ Qdrant nếu file pkl không tồn tại.""" print("BM25 index không tìm thấy, đang rebuild từ Qdrant...") client = get_client() texts, metadatas = [], [] offset = None while True: result, next_offset = client.scroll( collection_name=COLLECTION_NAME, limit=100, offset=offset, with_payload=True, with_vectors=False, ) for p in result: texts.append(p.payload.get("text", "")) metadatas.append({k: v for k, v in p.payload.items() if k != "text"}) if next_offset is None: break offset = next_offset tokenized = [_tokenize_vi(t) for t in texts] bm25 = BM25Okapi(tokenized) data = {"bm25": bm25, "texts": texts, "metadatas": metadatas} with open(BM25_PATH, "wb") as f: pickle.dump(data, f) print(f"Rebuild xong: {len(texts)} chunks.") return data def _load_bm25(): if not Path(BM25_PATH).exists(): return _build_bm25_from_qdrant() with open(BM25_PATH, "rb") as f: return pickle.load(f) def _reciprocal_rank_fusion(rankings: list[list[int]], k: int = 60) -> list[tuple[int, float]]: """Kết hợp nhiều bảng xếp hạng bằng RRF.""" scores = {} for ranking in rankings: for rank, doc_id in enumerate(ranking): scores[doc_id] = scores.get(doc_id, 0) + 1 / (k + rank + 1) return sorted(scores.items(), key=lambda x: x[1], reverse=True) def query(query_embedding: list[float], query_text: str, top_k: int = 5) -> list[dict]: """Hybrid retrieval: vector + BM25 + RRF.""" client = get_client() top_n = top_k * 4 # lấy nhiều hơn để RRF có đủ candidates # Vector search vector_results = client.query_points( collection_name=COLLECTION_NAME, query=query_embedding, limit=top_n, ).points vector_ids = [r.id for r in vector_results] vector_payload = {r.id: r.payload for r in vector_results} # BM25 search bm25_data = _load_bm25() bm25 = bm25_data["bm25"] texts = bm25_data["texts"] metadatas = bm25_data["metadatas"] tokenized_query = _tokenize_vi(query_text) bm25_scores = bm25.get_scores(tokenized_query) bm25_ids = sorted(range(len(bm25_scores)), key=lambda i: bm25_scores[i], reverse=True)[:top_n] # RRF fused = _reciprocal_rank_fusion([vector_ids, bm25_ids])[:top_k] # Build kết quả hits = [] for doc_id, rrf_score in fused: if doc_id in vector_payload: payload = vector_payload[doc_id] hits.append({ "text": payload["text"], "metadata": {k: v for k, v in payload.items() if k != "text"}, "score": round(rrf_score, 4), }) elif doc_id < len(texts): hits.append({ "text": texts[doc_id], "metadata": metadatas[doc_id], "score": round(rrf_score, 4), }) return hits