Cognitive-rag / run_pipeline.py
run_pipeline.py
Raw
#!/usr/bin/env python
"""
End-to-End RAG Pipeline Demo

Runs the full flow:
1. User query -> Embed
2. Retrieve (documents + conversation memory) using selectable strategy
3. Generate answer (Claude via Bedrock)
4. Store output as conversation memory

Supports retrieval strategies: dense, sparse, hybrid (with RRF)


python run_pipeline.py                    Uses dense (default)
python run_pipeline.py --strategy sparse  Uses sparse only
python run_pipeline.py --strategy hybrid  Uses hybrid with RRF
"""

import sys
import argparse
from pathlib import Path

# Load AWS credentials from .env
from dotenv import load_dotenv
load_dotenv(Path(__file__).parent / "05_generation" / ".env")

# Add all module paths - ORDER MATTERS for config.py resolution
PROJECT_ROOT = Path(__file__).parent

# Clear any cached config module to ensure we get the right one
if 'config' in sys.modules:
    del sys.modules['config']

# Add vectordb FIRST
sys.path.insert(0, str(PROJECT_ROOT / "04_vectoredb"))

# Import Pinecone modules BEFORE adding other paths
from pinecone_client import get_client, get_sparse_client
from retriever_multi import (
    get_unified_context_for_llm,
    export_retrieval_to_json,
    format_retrieval_output,
    calculate_retrieval_metrics,
)

# Now add other paths
sys.path.insert(0, str(PROJECT_ROOT / "03_embedding"))
sys.path.insert(0, str(PROJECT_ROOT / "05_generation"))
sys.path.insert(0, str(PROJECT_ROOT / "06_conversation_memory"))

from embed_query import load_client
from multihop_agent import MultiHopAgent, MultiHopConfig
from memory_processor import store_memory_sync, get_or_create_session_id

# Import config for NVIDIA API key
import config

# Set up LangSmith tracing
import os
os.environ["LANGCHAIN_TRACING_V2"] = config.LANGCHAIN_TRACING_V2
os.environ["LANGCHAIN_API_KEY"] = config.LANGCHAIN_API_KEY
os.environ["LANGCHAIN_PROJECT"] = config.LANGCHAIN_PROJECT

from langsmith import traceable


@traceable(name="embed_query")
def embed_query_nim(query: str, client):
    """Embed query using NVIDIA NIM embeddings."""
    embedding = client.embed_query(query)
    return embedding


@traceable(name="generate_answer")
def generate_answer_with_context(query: str, context_chunks: list):
    """
    Call LLM to generate answer using retrieved context.
    Uses Claude via Bedrock.
    """
    import json
    import boto3

    # Build context from retrieved chunks
    notes = "\n".join(f"- {chunk.get('text', '')}" for chunk in context_chunks)

    SYSTEM_PROMPT = """
You are a helpful assistant.

RULES:
- Use ONLY information explicitly stated in the provided context.
- Do NOT invent facts, dosages, or recommendations.
- If the context does not clearly answer the question, say:
  "The provided documents do not contain this information."

OUTPUT:
- Write a short natural-language answer (1-2 sentences).
"""

    prompt = f"""
Context:
{notes}

Question:
{query}

Answer based on the context above.
"""

    try:
        bedrock = boto3.client(
            service_name="bedrock-runtime",
            region_name="ca-central-1"
        )

        response = bedrock.invoke_model(
            modelId="anthropic.claude-3-haiku-20240307-v1:0",
            body=json.dumps({
                "anthropic_version": "bedrock-2023-05-31",
                "max_tokens": 300,
                "temperature": 0,
                "system": SYSTEM_PROMPT,
                "messages": [
                    {"role": "user", "content": prompt}
                ]
            }),
            contentType="application/json",
            accept="application/json"
        )

        response_body = json.loads(response["body"].read())
        answer = response_body["content"][0]["text"].strip()
        return answer

    except Exception as e:
        return f"[LLM Error: {e}]"


def format_sandy_output(query: str, answer: str, evidence_chunks: list) -> str:
    """Format output in Sandy's expected format for memory storage."""
    lines = []
    lines.append("Query:")
    lines.append(f'    "{query}"')
    lines.append("")
    lines.append("Answer:")
    lines.append(f"    {answer}")

    for chunk in evidence_chunks[:2]:
        lines.append("")
        lines.append("Evidence:")
        lines.append(f"    {chunk.get('text', '')[:200]}...")

    return "\n".join(lines)


def format_context_for_display(context: dict, max_chars: int = 1000) -> str:
    """Format retrieved context for display."""
    lines = []

    if context.get("memory_chunks"):
        lines.append("=== CONVERSATION MEMORY ===")
        for i, chunk in enumerate(context["memory_chunks"][:2], 1):
            score = chunk.get("original_score", chunk.get("score", 0))
            lines.append(f"[Memory {i}] (score: {score:.3f})")
            text = chunk.get("text", "")[:300]
            lines.append(text + "...")
            lines.append("")

    if context.get("document_chunks"):
        lines.append("=== DOCUMENTS ===")
        for i, chunk in enumerate(context["document_chunks"][:3], 1):
            score = chunk.get("cross_encoder_score", chunk.get("score", chunk.get("rrf_score", 0)))
            source = chunk.get("source", "dense")
            score_label = "ce_score" if "cross_encoder_score" in chunk else "score"
            lines.append(f"[Doc {i}] ({score_label}: {score:.4f}, source: {source})")
            text = chunk.get("text", "")[:300]
            lines.append(text + "...")
            lines.append("")

    output = "\n".join(lines)
    if len(output) > max_chars:
        output = output[:max_chars] + "\n..."
    return output


def _build_context(
    doc_chunks: list,
    memory_context: dict = None,
    related_chunks: list = None,
    extra: dict = None,
) -> dict:
    """
    Build a standardised context dict from retrieval results.

    Args:
        doc_chunks:      Retrieved document chunks.
        memory_context:  Output of get_unified_context_for_llm (for memory fields).
                         Pass None when memory retrieval was skipped.
        related_chunks:  List of related chunks fetched via fetch_related_chunks.
        extra:           Any additional keys to merge (e.g. dense_count, sparse_count).

    Returns:
        Context dict compatible with format_retrieval_output / run_pipeline display.
    """
    related_chunks = related_chunks or []
    mem = memory_context or {}

    ctx = {
        "document_chunks":    doc_chunks,
        "memory_chunks":      mem.get("memory_chunks", []),
        "related_chunks":     related_chunks,
        "doc_count":          len(doc_chunks),
        "memory_count":       mem.get("memory_count", 0),
        "related_count":      len(related_chunks),
        "memory_raw_count":   mem.get("memory_raw_count", 0),
        "memory_dropped_count": mem.get("memory_dropped_count", 0),
        "memory_min_score":   mem.get("memory_min_score", 0),
    }
    if extra:
        ctx.update(extra)
    return ctx


@traceable(name="retrieve_context")
def traced_retrieve(strategy: str, query: str, query_embedding: list, top_k: int = None, format_json: bool = True):
    """
    thin wrapper around get_unified_context_for_llm with tracing.

    args:
        strategy: "dense", "sparse", or "hybrid"
        query: raw query text
        query_embedding: embedded query vector
        top_k: number of docs to retrieve from Pinecone (default: config.TOP_K)
        format_json: if True, return JSON-formatted output; otherwise raw context dict
    """
    context = get_unified_context_for_llm(
        query_embedding=query_embedding,
        query_text=query,
        strategy=strategy,
        doc_top_k=top_k,
        include_related=True,
    )

    if format_json:
        return format_retrieval_output(query, context)
    return context


@traceable(name="rag_pipeline")
def run_pipeline(
    strategy: str = "dense",
    use_multihop: bool = True,
    multihop_config: MultiHopConfig | None = None,
):
    """Run the full RAG pipeline interactively."""
    print("\n" + "=" * 70)
    print(" CognitiveRAG End-to-End Pipeline")
    print(f" Retrieval Strategy: {strategy.upper()}")
    if config.RERANKER_ENABLED:
        print(f" Reranker: {config.RERANKER_MODEL} (top_n={config.RERANKER_TOP_N})")
    else:
        print(" Reranker: disabled")
    print("=" * 70)

    # Load NVIDIA NIM embedding client
    print("\n Loading NVIDIA NIM embedding model...")
    nim_client = load_client(config.EMBEDDING_MODEL, config.NVIDIA_API_KEY, "NONE")

    # Get Pinecone clients
    print(" Connecting to Pinecone...")
    client = get_client()
    index = client.get_index()

    # Show sparse index status if using sparse/hybrid
    if strategy in ["sparse", "hybrid"]:
        sparse_client = get_sparse_client()
        sparse_stats = sparse_client.get_index_stats()
        sparse_count = sparse_stats.get("total_vector_count", 0)
        print(f" Sparse index: {config.SPARSE_INDEX_NAME} ({sparse_count} vectors)")
        if sparse_count == 0:
            print(" [WARNING] Sparse index is empty. Run: python upsert_db.py --mode sparse")

    # Get session
    session_id = get_or_create_session_id()
    print(f" Session: {session_id[:8]}...")

    turn_number = 0

    while True:
        print("\n" + "-" * 70)
        query = input(" Enter your query (or 'quit' to exit): ").strip()

        if query.lower() in ['quit', 'exit', 'q']:
            print("\n Goodbye!")
            break

        if not query:
            continue

        turn_number += 1

        if use_multihop:
            print("\n Step 1: Running multi-hop agent...")
            agent_config = multihop_config or MultiHopConfig(
                strategy=strategy,
                per_hop_top_k=config.TOP_K,
            )
            agent = MultiHopAgent(
                embed_fn=lambda text: embed_query_nim(text, nim_client),
                config=agent_config,
            )
            result = agent.run(query)
            answer = result.answer
            evidence_chunks = [
                {"text": item.get("text", ""), "doc_id": item.get("doc_id", ""), "chunk_id": item.get("chunk_id", "")}
                for item in getattr(result, "evidence", [])
            ]
            print(f"    Termination: {result.termination_reason}")
            print(f"    Confidence: {result.confidence:.2f}")
            print(f"    Trace: {result.trace_path}")
        else:
            # Step 1: Embed query
            print("\n Step 1: Embedding query...")
            query_embedding = embed_query_nim(query, nim_client)

            # Step 2: Retrieve context based on strategy (traced)
            print(f" Step 2: Retrieving context ({strategy})...")
            context = traced_retrieve(strategy, query, query_embedding, top_k=config.TOP_K, format_json=False)
            doc_chunks = context["document_chunks"]

            # Display reranking info
            if context.get("reranked"):
                print(f"    Reranked with cross-encoder -> top {len(doc_chunks)} passages")

            # Calculate and display metrics
            metrics = calculate_retrieval_metrics(doc_chunks, top_k=config.TOP_K)

            print(f"    Documents: {context['doc_count']}")
            if strategy == "hybrid":
                print(f"      - Dense: {context.get('dense_count', 'N/A')}")
                print(f"      - Sparse: {context.get('sparse_count', 'N/A')}")
            print(f"    Related: {context.get('related_count', 0)}")

            # Memory stats
            if context['memory_raw_count'] == 0:
                print("    Memories: 0 (no memories in namespace)")
            elif context['memory_count'] == 0:
                print(f"    Memories: 0 (dropped {context['memory_dropped_count']} below threshold)")
            else:
                dropped_msg = f", dropped {context['memory_dropped_count']}" if context['memory_dropped_count'] > 0 else ""
                print(f"    Memories: {context['memory_count']}{dropped_msg}")

            # Show metrics
            print(f"    Metrics: avg={metrics['avg_score']:.3f}, var={metrics['score_variance']:.4f}, coverage={metrics['coverage_ratio']:.1%}")

            # Export retrieval to JSON
            output_json = export_retrieval_to_json(query, context)
            print(f"    Saved to: 04_vectoredb/output/retrieval_output.json")

            # Show formatted context
            print("\n" + "=" * 50)
            print("CONTEXT SENT TO LLM:")
            print("=" * 50)
            print(format_context_for_display(context))

            # Step 3: Generate answer
            print("\n Step 3: Calling LLM (Claude)...")
            all_chunks = context["document_chunks"] + context.get("memory_chunks", [])
            answer = generate_answer_with_context(query, all_chunks)
            evidence_chunks = context["document_chunks"]

        print("\n" + "=" * 50)
        print(" ANSWER:")
        print("=" * 50)
        print(answer)

        # Step 4: Store as memory
        print("\n Step 4: Storing as conversation memory...")
        sandy_output = format_sandy_output(query, answer, evidence_chunks)

        result = store_memory_sync(sandy_output, index, session_id, turn_number)
        if result:
            print(f"    Stored: {result['chunk_id']}")
        else:
            print("    Skipped (duplicate or error)")


def main():
    parser = argparse.ArgumentParser(description="CognitiveRAG Pipeline")
    parser.add_argument(
        "--strategy",
        choices=["dense", "sparse", "hybrid"],
        default="dense",
        help="Retrieval strategy (default: dense)"
    )
    parser.add_argument(
        "--single-hop",
        action="store_true",
        help="Use legacy single-hop retrieval instead of multi-hop agent.",
    )
    parser.add_argument("--max-hops", type=int, default=3)
    parser.add_argument("--per-hop-top-k", type=int, default=8)
    parser.add_argument("--total-doc-budget", type=int, default=24)
    parser.add_argument("--sufficiency-threshold", type=float, default=0.75)
    parser.add_argument("--time-limit-sec", type=int, default=60)
    args = parser.parse_args()

    multihop_config = MultiHopConfig(
        max_hops=args.max_hops,
        per_hop_top_k=args.per_hop_top_k,
        total_doc_budget=args.total_doc_budget,
        sufficiency_threshold=args.sufficiency_threshold,
        time_limit_sec=args.time_limit_sec,
        strategy=args.strategy,
    )

    run_pipeline(
        strategy=args.strategy,
        use_multihop=not args.single_hop,
        multihop_config=multihop_config,
    )


if __name__ == "__main__":
    main()