"""
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