#!/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()