Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 20 additions & 1 deletion apps/rag-pipeline/rag/__init__.py
Original file line number Diff line number Diff line change
@@ -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",
]
89 changes: 89 additions & 0 deletions apps/rag-pipeline/rag/pipeline.py
Original file line number Diff line number Diff line change
@@ -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))
81 changes: 81 additions & 0 deletions apps/rag-pipeline/tests/test_pipeline.py
Original file line number Diff line number Diff line change
@@ -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()
Loading