from typing import List, Dict from .config import get_settings from .gemini_client import GeminiClient from loguru import logger import asyncio class Reranker: def __init__(self): settings = get_settings() self.provider = getattr(settings, 'rerank_provider', settings.llm_provider) self.model = getattr(settings, 'rerank_model', settings.llm_model) if self.provider == 'gemini': self.client = GeminiClient() # elif self.provider == 'openai': # self.client = OpenAIClient(settings.openai_api_key, model=self.model) # elif self.provider == 'cohere': # self.client = CohereClient(settings.cohere_api_key, model=self.model) else: raise NotImplementedError(f"Rerank provider {self.provider} not supported yet.") async def _score_doc(self, query: str, doc: Dict) -> Dict: """ Score một document với query. """ content = (doc.get('tieude', '') or '') + ' ' + (doc.get('noidung', '') or '') prompt = ( f"Đoạn luật: {content}\n" f"Câu hỏi: {query}\n" "Hãy đánh giá mức độ liên quan giữa đoạn luật và câu hỏi trên thang điểm 0-10. " "Chỉ trả về một số duy nhất." ) try: if self.provider == 'gemini': loop = asyncio.get_event_loop() logger.info(f"[RERANK] Sending prompt to Gemini: {prompt}") score = await loop.run_in_executor(None, self.client.generate_text, prompt) logger.info(f"[RERANK] Got score from Gemini: {score}") else: raise NotImplementedError(f"Rerank provider {self.provider} not supported yet in rerank method.") score = float(str(score).strip().split()[0]) doc['rerank_score'] = score return doc except Exception as e: logger.error(f"[RERANK] Lỗi khi tính score: {e} | doc: {doc}") doc['rerank_score'] = 0 return doc async def rerank(self, query: str, docs: List[Dict], top_k: int = 5) -> List[Dict]: """ Rerank docs theo độ liên quan với query, trả về top_k docs. Sử dụng concurrency để process nhiều docs cùng lúc. """ logger.info(f"[RERANK] Start rerank for query: {query} | docs: {len(docs)} | top_k: {top_k}") if not docs: return [] # Giới hạn số docs để rerank (tối đa 10 docs) docs_to_rerank = docs[:10] if len(docs) > 10 else docs logger.info(f"[RERANK] Will rerank {len(docs_to_rerank)} docs (limited from {len(docs)})") # Process docs với concurrency batch_size = 5 # Process 5 docs cùng lúc scored = [] for i in range(0, len(docs_to_rerank), batch_size): batch = docs_to_rerank[i:i + batch_size] logger.info(f"[RERANK] Processing batch {i//batch_size + 1}: {len(batch)} docs") # Tạo tasks cho batch hiện tại tasks = [self._score_doc(query, doc) for doc in batch] # Chạy batch concurrently batch_results = await asyncio.gather(*tasks, return_exceptions=True) # Xử lý kết quả for result in batch_results: if isinstance(result, Exception): logger.error(f"[RERANK] Batch processing error: {result}") continue scored.append(result) logger.info(f"[RERANK] Completed batch {i//batch_size + 1}, processed {len(scored)} docs so far") # Sort theo score và trả về top_k scored = sorted(scored, key=lambda x: x['rerank_score'], reverse=True) result = scored[:top_k] logger.info(f"[RERANK] Top reranked docs: {result}") return result