#!/usr/bin/env python """ Multi-hop retrieval agent for CognitiveRAG. Coordinates query decomposition, iterative retrieval, evidence validation, and final answer synthesis with provenance-aware citations. """ from __future__ import annotations import json import os import re import time import uuid from dataclasses import dataclass, field from datetime import datetime from typing import Any, Dict, List, Optional, Tuple import boto3 from langsmith import traceable try: from retriever_multi import get_unified_context_for_llm except ImportError: import sys from pathlib import Path PROJECT_ROOT = Path(__file__).resolve().parents[1] sys.path.insert(0, str(PROJECT_ROOT / "04_vectoredb")) from retriever_multi import get_unified_context_for_llm @dataclass class MultiHopConfig: max_hops: int = 3 per_hop_top_k: int = 8 total_doc_budget: int = 24 sufficiency_threshold: float = 0.75 time_limit_sec: int = 60 strategy: str = "hybrid" include_related: bool = False memory_top_k: int = 3 @dataclass class EvidenceItem: chunk_id: str doc_id: str text: str source: str modality: str section: str page: Optional[int] score: float location: str def provenance_key(self) -> str: if self.chunk_id: return self.chunk_id base = f"{self.doc_id}|{self.page}|{self.section}|{self.location}" return f"{base}|{hash(self.text)}" @dataclass class HopRecord: hop_index: int sub_question: str retrieval_strategy: str retrieved_count: int new_evidence_count: int total_evidence_count: int validation: Dict[str, Any] retrieval_metadata: List[Dict[str, Any]] started_at: str completed_at: str @dataclass class MultiHopState: query: str started_at: str hops: List[HopRecord] = field(default_factory=list) evidence_pool: List[EvidenceItem] = field(default_factory=list) termination_reason: str = "" status: str = "running" completed_at: Optional[str] = None def to_dict(self) -> Dict[str, Any]: return { "query": self.query, "started_at": self.started_at, "completed_at": self.completed_at, "status": self.status, "termination_reason": self.termination_reason, "evidence_pool": [item.__dict__ for item in self.evidence_pool], "hops": [ { "hop_index": hop.hop_index, "sub_question": hop.sub_question, "retrieval_strategy": hop.retrieval_strategy, "retrieved_count": hop.retrieved_count, "new_evidence_count": hop.new_evidence_count, "total_evidence_count": hop.total_evidence_count, "validation": hop.validation, "retrieval_metadata": hop.retrieval_metadata, "started_at": hop.started_at, "completed_at": hop.completed_at, } for hop in self.hops ], } @dataclass class MultiHopResult: answer: str citations: List[Dict[str, Any]] confidence: float status: str termination_reason: str trace_path: str reasoning_trace: str = "" evidence: List[Dict[str, Any]] = field(default_factory=list) class MultiHopAgent: def __init__( self, embed_fn, config: MultiHopConfig, bedrock_region: str = "ca-central-1", output_dir: Optional[str] = None, ) -> None: self.embed_fn = embed_fn self.config = config self.bedrock = boto3.client( service_name="bedrock-runtime", region_name=bedrock_region, ) self.output_dir = output_dir or os.path.join( os.path.dirname(__file__), "output", "multihop_traces" ) os.makedirs(self.output_dir, exist_ok=True) @traceable(name="multihop_decompose") def decompose(self, query: str, evidence_summary: str) -> Dict[str, Any]: system_prompt = ( "You are a query decomposition agent. " "Respond with JSON only." ) user_prompt = f""" Query: {query} Evidence Summary: {evidence_summary} If the query can be answered with current evidence, return: {{"action":"answer","sub_question":""}} If more evidence is needed, return: {{"action":"ask","sub_question":""}} """ response = self._invoke_claude(system_prompt, user_prompt) payload = self._parse_json_response(response) if not payload: return {"action": "ask", "sub_question": query} action = payload.get("action", "ask") sub_question = payload.get("sub_question", "").strip() if action == "answer": return {"action": "answer", "sub_question": ""} return {"action": "ask", "sub_question": sub_question or query} @traceable(name="multihop_retrieve") def retrieve(self, sub_question: str) -> Tuple[List[Dict[str, Any]], Dict[str, Any]]: query_embedding = None if self.config.strategy in ("dense", "hybrid") or self.config.memory_top_k > 0: query_embedding = self.embed_fn(sub_question) context = get_unified_context_for_llm( query_embedding=query_embedding, query_text=sub_question, strategy=self.config.strategy, doc_top_k=self.config.per_hop_top_k, memory_top_k=self.config.memory_top_k, include_related=self.config.include_related, ) doc_chunks = context.get("document_chunks", []) memory_chunks = context.get("memory_chunks", []) return doc_chunks + memory_chunks, context def _dedupe_evidence( self, existing: Dict[str, EvidenceItem], candidates: List[EvidenceItem] ) -> Tuple[List[EvidenceItem], int]: new_items = [] for item in candidates: key = item.provenance_key() if key not in existing: existing[key] = item new_items.append(item) return new_items, len(new_items) def _build_evidence_item(self, chunk: Dict[str, Any]) -> EvidenceItem: doc_id = str(chunk.get("doc_id") or chunk.get("session_id") or "") chunk_id = str(chunk.get("chunk_id") or "") modality = str(chunk.get("modality") or "text") section = str(chunk.get("section") or "") page = chunk.get("page") source = str(chunk.get("source") or "document") score = float( chunk.get("cross_encoder_score") or chunk.get("score") or chunk.get("rrf_score") or 0.0 ) location = "" if page is not None: location = f"page:{page}" elif section: location = f"section:{section}" else: location = "unknown" return EvidenceItem( chunk_id=chunk_id, doc_id=doc_id, text=str(chunk.get("text") or ""), source=source, modality=modality, section=section, page=page, score=score, location=location, ) def _summarize_evidence(self, evidence: List[EvidenceItem], max_chars: int = 1200) -> str: if not evidence: return "No evidence collected yet." lines = [] for item in evidence[:6]: snippet = item.text.replace("\n", " ").strip()[:180] lines.append( f"- [{item.doc_id or 'unknown'}|{item.location}] {snippet}" ) summary = "\n".join(lines) if len(summary) > max_chars: return summary[:max_chars] + "..." return summary def _heuristic_validation(self, query: str, evidence: List[EvidenceItem]) -> Dict[str, float]: if not evidence: return {"keyword_coverage": 0.0, "doc_diversity": 0.0, "heuristic_score": 0.0} query_tokens = [ token for token in re.findall(r"[A-Za-z0-9]{4,}", query.lower()) ] if not query_tokens: return {"keyword_coverage": 0.0, "doc_diversity": 0.0, "heuristic_score": 0.0} evidence_text = " ".join(item.text.lower() for item in evidence) matches = sum(1 for token in set(query_tokens) if token in evidence_text) keyword_coverage = matches / max(len(set(query_tokens)), 1) unique_docs = len(set(item.doc_id for item in evidence if item.doc_id)) doc_diversity = min(unique_docs / 3.0, 1.0) heuristic_score = round((0.6 * keyword_coverage) + (0.4 * doc_diversity), 4) return { "keyword_coverage": round(keyword_coverage, 4), "doc_diversity": round(doc_diversity, 4), "heuristic_score": heuristic_score, } @traceable(name="multihop_validate") def validate(self, query: str, evidence: List[EvidenceItem]) -> Dict[str, Any]: heuristic = self._heuristic_validation(query, evidence) summary = self._summarize_evidence(evidence) system_prompt = "You are a strict evidence sufficiency critic. Respond with JSON only." user_prompt = f""" Question: {query} Evidence Summary: {summary} Return JSON with: {{"sufficiency_score":0-1,"relevance_score":0-1,"missing_info":"...","should_stop":true/false}} """ response = self._invoke_claude(system_prompt, user_prompt) payload = self._parse_json_response(response) or {} sufficiency = float(payload.get("sufficiency_score", 0.0)) relevance = float(payload.get("relevance_score", 0.0)) should_stop = bool(payload.get("should_stop", False)) combined = (0.55 * sufficiency) + (0.25 * relevance) + (0.20 * heuristic["heuristic_score"]) return { "heuristic": heuristic, "sufficiency_score": round(sufficiency, 4), "relevance_score": round(relevance, 4), "combined_score": round(combined, 4), "missing_info": payload.get("missing_info", ""), "should_stop": should_stop, } @traceable(name="multihop_synthesize") def synthesize(self, query: str, evidence: List[EvidenceItem]) -> Dict[str, Any]: evidence_payload = [ { "chunk_id": item.chunk_id, "doc_id": item.doc_id, "source": item.source, "location": item.location, "text": item.text, } for item in evidence ] system_prompt = ( "You are a grounded QA system. " "Use only the provided evidence. " "Respond with JSON only." ) user_prompt = f""" Question: {query} Evidence: {json.dumps(evidence_payload[:12], ensure_ascii=True, indent=2)} Return JSON: {{"answer":"...","citations":[{{"chunk_id":"","doc_id":"","location":""}}],"confidence":0-1,"reasoning_trace":"short"}} If evidence is insufficient, say so directly in the answer and set confidence <= 0.3. """ response = self._invoke_claude(system_prompt, user_prompt) payload = self._parse_json_response(response) or {} return { "answer": payload.get("answer", "The provided evidence is insufficient."), "citations": payload.get("citations", []), "confidence": float(payload.get("confidence", 0.0) or 0.0), "reasoning_trace": payload.get("reasoning_trace", ""), } def _invoke_claude(self, system_prompt: str, user_prompt: str) -> str: body = json.dumps( { "anthropic_version": "bedrock-2023-05-31", "max_tokens": 500, "temperature": 0, "system": system_prompt, "messages": [{"role": "user", "content": user_prompt}], } ) response = self.bedrock.invoke_model( modelId="anthropic.claude-3-haiku-20240307-v1:0", body=body, contentType="application/json", accept="application/json", ) response_body = json.loads(response["body"].read()) return response_body["content"][0]["text"].strip() def _parse_json_response(self, text: str) -> Optional[Dict[str, Any]]: try: return json.loads(text) except json.JSONDecodeError: match = re.search(r"\{.*\}", text, re.DOTALL) if not match: return None try: return json.loads(match.group(0)) except json.JSONDecodeError: return None def _write_trace(self, state: MultiHopState, final: Dict[str, Any]) -> str: trace_id = uuid.uuid4().hex[:8] timestamp = datetime.utcnow().strftime("%Y%m%d_%H%M%S") path = os.path.join(self.output_dir, f"multihop_trace_{timestamp}_{trace_id}.json") payload = { "state": state.to_dict(), "final": final, } with open(path, "w", encoding="utf-8") as handle: json.dump(payload, handle, indent=2, ensure_ascii=True) return path @traceable(name="multihop_run") def run(self, query: str) -> MultiHopResult: start_time = time.monotonic() state = MultiHopState( query=query, started_at=datetime.utcnow().isoformat(timespec="seconds") + "Z", ) evidence_index: Dict[str, EvidenceItem] = {} retrieved_total = 0 for hop in range(1, self.config.max_hops + 1): hop_start = datetime.utcnow().isoformat(timespec="seconds") + "Z" elapsed = time.monotonic() - start_time if elapsed > self.config.time_limit_sec: state.status = "terminated" state.termination_reason = "time_limit_exceeded" break evidence_summary = self._summarize_evidence(state.evidence_pool) decomposition = self.decompose(query, evidence_summary) if decomposition["action"] == "answer": state.status = "terminated" state.termination_reason = "decomposer_complete" break sub_question = decomposition["sub_question"] chunks, context = self.retrieve(sub_question) retrieved_total += len(chunks) candidates = [self._build_evidence_item(chunk) for chunk in chunks] new_items, new_count = self._dedupe_evidence(evidence_index, candidates) state.evidence_pool.extend(new_items) validation = self.validate(query, state.evidence_pool) should_stop = ( validation["combined_score"] >= self.config.sufficiency_threshold or validation.get("should_stop", False) ) hop_record = HopRecord( hop_index=hop, sub_question=sub_question, retrieval_strategy=context.get("strategy", self.config.strategy), retrieved_count=len(chunks), new_evidence_count=new_count, total_evidence_count=len(state.evidence_pool), validation=validation, retrieval_metadata=[ { "chunk_id": c.get("chunk_id", ""), "doc_id": c.get("doc_id", ""), "section": c.get("section", ""), "page": c.get("page"), "source": c.get("source", ""), "score": c.get("score", c.get("rrf_score", 0)), } for c in chunks ], started_at=hop_start, completed_at=datetime.utcnow().isoformat(timespec="seconds") + "Z", ) state.hops.append(hop_record) if retrieved_total >= self.config.total_doc_budget: state.status = "terminated" state.termination_reason = "retrieval_budget_exceeded" break if should_stop: state.status = "terminated" state.termination_reason = "evidence_sufficient" break if state.status != "terminated": state.status = "terminated" if not state.termination_reason: state.termination_reason = "max_hops_reached" synthesis = self.synthesize(query, state.evidence_pool) state.completed_at = datetime.utcnow().isoformat(timespec="seconds") + "Z" trace_path = self._write_trace(state, synthesis) evidence_payload = [ { "chunk_id": item.chunk_id, "doc_id": item.doc_id, "text": item.text, "source": item.source, "location": item.location, } for item in state.evidence_pool ] return MultiHopResult( answer=synthesis["answer"], citations=synthesis.get("citations", []), confidence=synthesis.get("confidence", 0.0), reasoning_trace=synthesis.get("reasoning_trace", ""), status=state.status, termination_reason=state.termination_reason, trace_path=trace_path, evidence=evidence_payload, )