""" Multi-Namespace Retriever for CognitiveRAG. Supports modular retrieval strategies (Dense, Sparse, Hybrid+RRF) for Selection Agent. """ from __future__ import annotations from enum import Enum from typing import Dict, List, Any, Optional import statistics try: from .pinecone_client import get_client, get_sparse_client from . import config from .reranker import rerank except ImportError: from pinecone_client import get_client, get_sparse_client import config from reranker import rerank # Set up LangSmith tracing import os os.environ.setdefault("LANGCHAIN_TRACING_V2", config.LANGCHAIN_TRACING_V2) os.environ.setdefault("LANGCHAIN_API_KEY", config.LANGCHAIN_API_KEY) os.environ.setdefault("LANGCHAIN_PROJECT", config.LANGCHAIN_PROJECT) from langsmith import traceable class RetrievalStrategy(Enum): """Retrieval strategy for Selection Agent.""" DENSE = "dense" SPARSE = "sparse" HYBRID = "hybrid" def reciprocal_rank_fusion( results_list: List[List[Dict]], k: int = 60 ) -> List[Dict]: """ Merge multiple result lists using Reciprocal Rank Fusion. RRF score = sum(1 / (k + rank)) for each result across all lists. """ scores = {} chunk_data = {} for results in results_list: for rank, doc in enumerate(results): doc_id = doc["chunk_id"] scores[doc_id] = scores.get(doc_id, 0) + 1 / (k + rank + 1) if doc_id not in chunk_data: chunk_data[doc_id] = doc sorted_ids = sorted(scores.items(), key=lambda x: x[1], reverse=True) merged = [] for doc_id, rrf_score in sorted_ids: doc = chunk_data[doc_id].copy() doc["rrf_score"] = round(rrf_score, 6) merged.append(doc) return merged def _tag_and_wrap(results: Dict, strategy: str) -> Dict[str, Any]: """Tag retrieved chunks with their source strategy and wrap in standard return dict.""" chunks = results.get("retrieved_chunks", []) for chunk in chunks: chunk["source"] = strategy return {"retrieved_chunks": chunks, "strategy": strategy, "count": len(chunks)} @traceable(name="retrieve_dense") def retrieve_dense( query_embedding: List[float], top_k: int = None, namespace: str = "" ) -> Dict[str, Any]: """Pure dense retrieval from existing index.""" client = get_client() top_k = top_k or config.TOP_K results = client.retrieve_similar( query_embedding=query_embedding, top_k=top_k, namespace=namespace, ) return _tag_and_wrap(results, "dense") @traceable(name="retrieve_sparse") def retrieve_sparse( query_text: str, top_k: int = None, namespace: str = "" ) -> Dict[str, Any]: """Pure sparse retrieval from sparse index.""" sparse_client = get_sparse_client() top_k = top_k or config.TOP_K results = sparse_client.retrieve_sparse( query=query_text, top_k=top_k, namespace=namespace, ) return _tag_and_wrap(results, "sparse") @traceable(name="retrieve_hybrid") def retrieve_hybrid( query_embedding: List[float], query_text: str, top_k: int = None, rrf_k: int = None, dense_namespace: str = "", sparse_namespace: str = "", ) -> Dict[str, Any]: """hybrid retrieval: dense + sparse merged with RRF.""" top_k = top_k or config.TOP_K rrf_k = rrf_k or config.RRF_K dense_results = retrieve_dense(query_embedding, top_k=top_k, namespace=dense_namespace) sparse_results = retrieve_sparse(query_text, top_k=top_k, namespace=sparse_namespace) merged = reciprocal_rank_fusion( [dense_results["retrieved_chunks"], sparse_results["retrieved_chunks"]], k=rrf_k ) for chunk in merged: chunk["source"] = "hybrid" return { "retrieved_chunks": merged[:top_k], "strategy": "hybrid", "count": len(merged[:top_k]), "dense_count": dense_results["count"], "sparse_count": sparse_results["count"], } def retrieve( query_embedding: List[float] = None, query_text: str = None, strategy: RetrievalStrategy = RetrievalStrategy.DENSE, top_k: int = None, namespace: str = "", ) -> Dict[str, Any]: """ Unified retrieval interface for Selection Agent. Returns standardized output compatible with all strategies. """ if strategy == RetrievalStrategy.DENSE: if query_embedding is None: raise ValueError("query_embedding required for DENSE strategy") return retrieve_dense(query_embedding, top_k, namespace) elif strategy == RetrievalStrategy.SPARSE: if query_text is None: raise ValueError("query_text required for SPARSE strategy") return retrieve_sparse(query_text, top_k, namespace) elif strategy == RetrievalStrategy.HYBRID: if query_embedding is None or query_text is None: raise ValueError("Both query_embedding and query_text required for HYBRID strategy") return retrieve_hybrid(query_embedding, query_text, top_k) else: raise ValueError(f"Unknown strategy: {strategy}") def calculate_retrieval_metrics(chunks: List[Dict], top_k: int = 5) -> Dict[str, float]: """ Calculate deterministic metrics for retrieved chunks. Used by LangGraph metrics node. """ if not chunks: return { "avg_score": 0.0, "score_variance": 0.0, "coverage_ratio": 0.0, "min_score": 0.0, "max_score": 0.0, } scores = [c.get("score", 0) or c.get("rrf_score", 0) for c in chunks] return { "avg_score": round(statistics.mean(scores), 4), "score_variance": round(statistics.variance(scores), 4) if len(scores) > 1 else 0.0, "coverage_ratio": round(len(chunks) / top_k, 4), "min_score": round(min(scores), 4), "max_score": round(max(scores), 4), } @traceable(name="get_unified_context") def get_unified_context_for_llm( query_embedding: List[float], query_text: str = None, strategy: str = "dense", doc_top_k: int = None, memory_top_k: int = 3, memory_boost: float = None, memory_min_score: float = None, doc_namespace: str = "", memory_namespace: str = None, include_related: bool = False, related_types: List[str] = None, export_json: bool = True, json_output_path: str = None, ) -> Dict[str, Any]: """ single entry point for all retrieval strategies. retrieves document chunks (dense, sparse, or hybrid), conversation memories, and optionally related chunks. applies cross-encoder reranking when enabled. args: query_embedding: dense vector for the query. query_text: raw query string (required for sparse/hybrid/reranking). strategy: "dense", "sparse", or "hybrid". doc_top_k: how many docs to pull from Pinecone (defaults to config.TOP_K). memory_top_k: how many memories to pull. include_related: whether to fetch prev/next chunks. export_json: whether to export results to JSON (default True). json_output_path: custom path for the JSON file (optional). returns: dict with document_chunks, memory_chunks, related_chunks, unified_context, and stats. """ client = get_client() doc_top_k = doc_top_k or config.TOP_K memory_boost = memory_boost if memory_boost is not None else config.MEMORY_BOOST memory_min_score = memory_min_score if memory_min_score is not None else config.MEMORY_MIN_SCORE memory_namespace = memory_namespace or config.CONVERSATION_MEMORY_NAMESPACE related_types = related_types or ["prev", "next"] # ── step 1: retrieve document chunks based on strategy ────────── extra_stats = {} if strategy == "sparse": if query_text is None: raise ValueError("query_text required for sparse strategy") result = retrieve_sparse(query_text, top_k=doc_top_k, namespace=doc_namespace) doc_chunks = result["retrieved_chunks"] elif strategy == "hybrid": if query_text is None or query_embedding is None: raise ValueError("both query_embedding and query_text required for hybrid strategy") result = retrieve_hybrid( query_embedding, query_text, top_k=doc_top_k, dense_namespace=doc_namespace, sparse_namespace=doc_namespace, ) doc_chunks = result["retrieved_chunks"] extra_stats["dense_count"] = result.get("dense_count", 0) extra_stats["sparse_count"] = result.get("sparse_count", 0) else: # dense (default) raw = client.retrieve_similar( query_embedding=query_embedding, top_k=doc_top_k, namespace=doc_namespace, ) doc_chunks = raw.get("retrieved_chunks", []) for chunk in doc_chunks: chunk["source"] = "document" # ── step 2: cross-encoder reranking ───────────────────────────── reranked = False if config.RERANKER_ENABLED and query_text and doc_chunks: doc_chunks = rerank(query_text, doc_chunks, top_n=config.RERANKER_TOP_N) reranked = True # ── step 3: retrieve conversation memories ────────────────────── memory_results = client.retrieve_similar( query_embedding=query_embedding, top_k=memory_top_k, namespace=memory_namespace, ) raw_memory_chunks = memory_results.get("retrieved_chunks", []) memory_chunks = [] dropped_memory_count = 0 for chunk in raw_memory_chunks: chunk["source"] = "conversation_memory" chunk["original_score"] = chunk["score"] if chunk["score"] < memory_min_score: dropped_memory_count += 1 continue chunk["score"] = min(chunk["score"] + memory_boost, 1.0) memory_chunks.append(chunk) # ── step 4: fetch related (prev/next) chunks ──────────────────── related_chunks = [] if include_related and doc_chunks: related_chunks = fetch_related_chunks( doc_chunks, related_types, doc_namespace, client ) # ── step 5: merge and sort ────────────────────────────────────── unified = doc_chunks + memory_chunks + related_chunks unified = sorted( unified, key=lambda x: x.get("cross_encoder_score", x.get("score", 0)), reverse=True, ) context = { "query": query_text, "strategy": strategy + ("+rerank" if reranked else ""), "document_chunks": doc_chunks, "memory_chunks": memory_chunks, "related_chunks": related_chunks, "unified_context": unified, "doc_count": len(doc_chunks), "memory_count": len(memory_chunks), "related_count": len(related_chunks), "memory_raw_count": len(raw_memory_chunks), "memory_dropped_count": dropped_memory_count, "memory_min_score": memory_min_score, "reranked": reranked, **extra_stats, } # ── step 6: export to JSON by default ─────────────────────────── if export_json and query_text: export_retrieval_to_json(query_text, context, output_path=json_output_path) return context def fetch_related_chunks( chunks: List[Dict], related_types: List[str], namespace: str, client = None ) -> List[Dict]: """ Fetch related chunks based on relational chunking schema. """ if client is None: client = get_client() related_ids = set() for chunk in chunks: # Collect direct prev/next IDs from chunk metadata for key in ("prev", "next"): if key in related_types: ref_id = chunk.get(f"{key}_chunk") or "" if ref_id: related_ids.add(ref_id) # Collect from nested related_chunks dict related = chunk.get("related_chunks", {}) if isinstance(related, dict): for rel_type in related_types: rel_value = related.get(rel_type) if rel_value: if isinstance(rel_value, list): related_ids.update(rel_value) else: related_ids.add(rel_value) if not related_ids: return [] index = client.get_index() try: fetch_result = index.fetch( ids=list(related_ids), namespace=namespace ) related_chunks = [] for chunk_id, vector_data in fetch_result.get("vectors", {}).items(): metadata = vector_data.get("metadata", {}) related_chunks.append({ "chunk_id": chunk_id, "text": metadata.get("text", ""), "doc_id": metadata.get("doc_id", ""), "section": metadata.get("section", ""), "page": metadata.get("page"), "source": "related", }) return related_chunks except Exception as e: import logging logging.warning(f"Failed to fetch related chunks: {e}") return [] def format_retrieval_output( query_text: str, context: Dict, ) -> Dict: """Format retrieval results into a structured JSON output.""" def _texts(key: str) -> List[str]: return [c.get("text", "") for c in context.get(key, [])] doc_metadata = [ { "chunk_id": c.get("chunk_id", ""), "doc_id": c.get("doc_id", ""), "section": c.get("section", ""), "page": c.get("page"), "score": round(c.get("score", 0), 4), } for c in context.get("document_chunks", []) ] memory_metadata = [ { "chunk_id": c.get("chunk_id", ""), "session_id": c.get("doc_id", ""), "section": c.get("section", ""), "turn": c.get("page"), "score": round(c.get("score", 0), 4), "original_score": round(c.get("original_score", 0), 4), } for c in context.get("memory_chunks", []) ] related_metadata = [ { "chunk_id": c.get("chunk_id", ""), "doc_id": c.get("doc_id", ""), "section": c.get("section", ""), "page": c.get("page"), } for c in context.get("related_chunks", []) ] return { "query": query_text, "retrieved_chunks": _texts("document_chunks"), "chunk_metadata": doc_metadata, "memory_chunks": _texts("memory_chunks"), "memory_metadata": memory_metadata, "related_chunks": _texts("related_chunks"), "related_metadata": related_metadata, "stats": { "document_count": context.get("doc_count", 0), "memory_count": context.get("memory_count", 0), "related_count": context.get("related_count", 0), "memory_raw_count": context.get("memory_raw_count", 0), "memory_dropped_count": context.get("memory_dropped_count", 0), }, } def export_retrieval_to_json( query_text: str, context: Dict, output_path: str = None, ) -> Dict: """ Export retrieval results to a JSON file. """ import json import os output = format_retrieval_output(query_text, context) if output_path is None: output_dir = os.path.join(os.path.dirname(__file__), "output") os.makedirs(output_dir, exist_ok=True) output_path = os.path.join(output_dir, "retrieval_output.json") with open(output_path, "w", encoding="utf-8") as f: json.dump(output, f, indent=2, ensure_ascii=False) return output