From 320ded8a50f3f9aaf3a937784eadc85cef15c688 Mon Sep 17 00:00:00 2001 From: Amp Date: Fri, 28 Aug 2026 13:28:15 +0000 Subject: [PATCH] feat(rag): add provider-independent pipeline runtime Co-authored-by: Aditya Balakrishnan --- apps/rag-pipeline/rag/__init__.py | 21 +++++- apps/rag-pipeline/rag/pipeline.py | 89 ++++++++++++++++++++++++ apps/rag-pipeline/tests/test_pipeline.py | 81 +++++++++++++++++++++ 3 files changed, 190 insertions(+), 1 deletion(-) create mode 100644 apps/rag-pipeline/rag/pipeline.py create mode 100644 apps/rag-pipeline/tests/test_pipeline.py diff --git a/apps/rag-pipeline/rag/__init__.py b/apps/rag-pipeline/rag/__init__.py index bc22615..94b9ec9 100644 --- a/apps/rag-pipeline/rag/__init__.py +++ b/apps/rag-pipeline/rag/__init__.py @@ -1,5 +1,24 @@ """Pure Python RAG foundation utilities.""" from .core import reciprocal_rank_fusion, split_text +from .pipeline import ( + INSUFFICIENT_INFORMATION, + Citation, + Generator, + RAGPipeline, + RAGResult, + RetrievalHit, + Retriever, +) -__all__ = ["reciprocal_rank_fusion", "split_text"] +__all__ = [ + "INSUFFICIENT_INFORMATION", + "Citation", + "Generator", + "RAGPipeline", + "RAGResult", + "RetrievalHit", + "Retriever", + "reciprocal_rank_fusion", + "split_text", +] diff --git a/apps/rag-pipeline/rag/pipeline.py b/apps/rag-pipeline/rag/pipeline.py new file mode 100644 index 0000000..8d6fb55 --- /dev/null +++ b/apps/rag-pipeline/rag/pipeline.py @@ -0,0 +1,89 @@ +"""Provider-independent orchestration for retrieval-augmented generation.""" + +from collections.abc import Mapping, Sequence +from dataclasses import dataclass, field +from typing import Protocol + +from .core import reciprocal_rank_fusion + +INSUFFICIENT_INFORMATION = "Insufficient information in the retrieved sources." + + +@dataclass(frozen=True) +class RetrievalHit: + """A retriever result with an identity shared across retrieval methods.""" + + document_id: str + text: str + metadata: Mapping[str, object] = field(default_factory=dict) + + +class Retriever(Protocol): + def __call__(self, query: str, limit: int) -> Sequence[RetrievalHit]: ... + + +class Generator(Protocol): + def __call__(self, question: str, context: str) -> str: ... + + +@dataclass(frozen=True) +class Citation: + label: str + document_id: str + metadata: Mapping[str, object] + + +@dataclass(frozen=True) +class RAGResult: + answer: str + context: str + citations: tuple[Citation, ...] + insufficient_information: bool = False + + +class RAGPipeline: + """Fuse dense and lexical retrieval before invoking an injected generator.""" + + def __init__( + self, + dense_retriever: Retriever, + lexical_retriever: Retriever, + generator: Generator, + *, + retrieval_limit: int = 5, + ) -> None: + if retrieval_limit <= 0: + raise ValueError("retrieval_limit must be positive") + self._dense_retriever = dense_retriever + self._lexical_retriever = lexical_retriever + self._generator = generator + self._retrieval_limit = retrieval_limit + + def run(self, question: str) -> RAGResult: + rankings = [ + self._dense_retriever(question, self._retrieval_limit), + self._lexical_retriever(question, self._retrieval_limit), + ] + hits_by_id: dict[str, RetrievalHit] = {} + id_rankings: list[list[str]] = [] + for ranking in rankings: + ids: list[str] = [] + for hit in ranking: + hits_by_id.setdefault(hit.document_id, hit) + ids.append(hit.document_id) + id_rankings.append(ids) + + fused = reciprocal_rank_fusion(id_rankings, limit=self._retrieval_limit) + if not fused: + return RAGResult(INSUFFICIENT_INFORMATION, "", (), True) + + context_parts: list[str] = [] + citations: list[Citation] = [] + for index, (document_id, _score) in enumerate(fused, start=1): + hit = hits_by_id[document_id] + label = f"SOURCE_{index}" + context_parts.append(f"[{label}]\n{hit.text}") + citations.append(Citation(label, document_id, dict(hit.metadata))) + + context = "\n\n".join(context_parts) + return RAGResult(self._generator(question, context), context, tuple(citations)) diff --git a/apps/rag-pipeline/tests/test_pipeline.py b/apps/rag-pipeline/tests/test_pipeline.py new file mode 100644 index 0000000..d93e665 --- /dev/null +++ b/apps/rag-pipeline/tests/test_pipeline.py @@ -0,0 +1,81 @@ +import unittest + +from rag import Citation, RAGPipeline, RetrievalHit + + +class RAGPipelineTests(unittest.TestCase): + def test_fuses_retrievers_and_builds_stably_labeled_context(self) -> None: + calls: list[tuple[str, str]] = [] + + def dense(query: str, limit: int) -> list[RetrievalHit]: + self.assertEqual((query, limit), ("What is AgentX?", 3)) + return [ + RetrievalHit("dense", "Dense-only text", {"url": "dense.example"}), + RetrievalHit("shared", "Shared grounded text", {"page": 7}), + ] + + def lexical(query: str, limit: int) -> list[RetrievalHit]: + return [ + RetrievalHit("shared", "Shared grounded text", {"page": 7}), + RetrievalHit("lexical", "Lexical-only text", {"section": "intro"}), + ] + + def generate(question: str, context: str) -> str: + calls.append((question, context)) + return "AgentX is grounded [SOURCE_1]." + + result = RAGPipeline(dense, lexical, generate, retrieval_limit=3).run("What is AgentX?") + + self.assertFalse(result.insufficient_information) + self.assertEqual(result.answer, "AgentX is grounded [SOURCE_1].") + self.assertEqual( + result.context, + "[SOURCE_1]\nShared grounded text\n\n" + "[SOURCE_2]\nDense-only text\n\n" + "[SOURCE_3]\nLexical-only text", + ) + self.assertEqual(calls, [("What is AgentX?", result.context)]) + self.assertEqual( + result.citations, + ( + Citation("SOURCE_1", "shared", {"page": 7}), + Citation("SOURCE_2", "dense", {"url": "dense.example"}), + Citation("SOURCE_3", "lexical", {"section": "intro"}), + ), + ) + + def test_returns_explicit_insufficient_result_without_generation(self) -> None: + generated = False + + def retrieve(query: str, limit: int) -> list[RetrievalHit]: + return [] + + def generate(question: str, context: str) -> str: + nonlocal generated + generated = True + return "unsupported" + + result = RAGPipeline(retrieve, retrieve, generate).run("unknown") + + self.assertTrue(result.insufficient_information) + self.assertEqual(result.answer, "Insufficient information in the retrieved sources.") + self.assertEqual(result.context, "") + self.assertEqual(result.citations, ()) + self.assertFalse(generated) + + def test_deduplicates_by_id_and_preserves_first_seen_metadata(self) -> None: + dense_hit = RetrievalHit("same", "canonical text", {"provider": "dense"}) + lexical_hit = RetrievalHit("same", "other text", {"provider": "lexical"}) + + result = RAGPipeline( + lambda _query, _limit: [dense_hit], + lambda _query, _limit: [lexical_hit], + lambda _question, _context: "answer", + ).run("question") + + self.assertEqual(result.context, "[SOURCE_1]\ncanonical text") + self.assertEqual(result.citations[0].metadata, {"provider": "dense"}) + + +if __name__ == "__main__": + unittest.main()