Cognitive-rag / 04_vectoredb / retriever_multi.py
retriever_multi.py
Raw
"""
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