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
3 changes: 3 additions & 0 deletions apps/rag-pipeline/rag/__init__.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
"""Pure Python RAG foundation utilities."""

from .core import reciprocal_rank_fusion, split_text
from .local_index import Document, LocalDocumentIndex
from .pipeline import (
INSUFFICIENT_INFORMATION,
Citation,
Expand All @@ -14,7 +15,9 @@
__all__ = [
"INSUFFICIENT_INFORMATION",
"Citation",
"Document",
"Generator",
"LocalDocumentIndex",
"RAGPipeline",
"RAGResult",
"RetrievalHit",
Expand Down
66 changes: 66 additions & 0 deletions apps/rag-pipeline/rag/local_index.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,66 @@
"""Deterministic in-memory document chunking and lexical retrieval."""

import re
from collections.abc import Iterable, Mapping
from dataclasses import dataclass, field

from .core import split_text
from .pipeline import RetrievalHit


def _tokens(text: str) -> set[str]:
return set(re.findall(r"\b\w+\b", text.casefold()))


@dataclass(frozen=True)
class Document:
document_id: str
text: str
metadata: Mapping[str, object] = field(default_factory=dict)


class LocalDocumentIndex:
"""A rebuildable lexical index compatible with ``RAGPipeline`` retrievers."""

def __init__(
self,
*,
chunk_size: int = 1200,
chunk_overlap: int = 300,
min_chunk_chars: int = 40,
) -> None:
self._chunk_options = {
"chunk_size": chunk_size,
"chunk_overlap": chunk_overlap,
"min_chunk_chars": min_chunk_chars,
}
self._chunks: list[tuple[RetrievalHit, set[str]]] = []

def rebuild(self, documents: Iterable[Document]) -> None:
"""Replace the complete index with caller-supplied documents."""
rebuilt: list[tuple[RetrievalHit, set[str]]] = []
seen: set[str] = set()
for document in documents:
if document.document_id in seen:
raise ValueError(f"duplicate document id: {document.document_id}")
seen.add(document.document_id)
metadata = dict(document.metadata)
metadata["document_id"] = document.document_id
for index, chunk in enumerate(split_text(document.text, **self._chunk_options)):
chunk_id = f"{len(document.document_id)}:{document.document_id}:{index}"
rebuilt.append((RetrievalHit(chunk_id, chunk, metadata), _tokens(chunk)))
self._chunks = rebuilt

def __call__(self, query: str, limit: int) -> list[RetrievalHit]:
if limit <= 0:
return []
query_tokens = _tokens(query)
if not query_tokens:
return []
ranked = [
(len(query_tokens & tokens), hit)
for hit, tokens in self._chunks
if query_tokens & tokens
]
ranked.sort(key=lambda item: (-item[0], item[1].document_id))
return [hit for _score, hit in ranked[:limit]]
62 changes: 62 additions & 0 deletions apps/rag-pipeline/tests/test_local_index.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,62 @@
import unittest

from rag import Document, LocalDocumentIndex


class LocalDocumentIndexTests(unittest.TestCase):
def test_indexes_caller_documents_without_title_or_filename_collisions(self) -> None:
index = LocalDocumentIndex(chunk_size=30, chunk_overlap=0, min_chunk_chars=1)
index.rebuild(
[
Document("doc/a", "shared title alpha unique", {"source": "first", "title": "Same"}),
Document("doc:b", "shared title beta unique", {"source": "second", "title": "Same"}),
]
)

alpha = index("alpha", 5)
beta = index("beta", 5)

self.assertEqual(len(alpha), 1)
self.assertEqual(len(beta), 1)
self.assertNotEqual(alpha[0].document_id, beta[0].document_id)
self.assertEqual(alpha[0].metadata, {"source": "first", "title": "Same", "document_id": "doc/a"})
self.assertEqual(beta[0].metadata["document_id"], "doc:b")

def test_rebuild_replaces_previous_contents_and_is_idempotent(self) -> None:
index = LocalDocumentIndex(min_chunk_chars=1)
documents = [Document("stable", "repeatable needle", {"source": "caller"})]

index.rebuild(documents)
first = index("needle", 10)
index.rebuild(documents)
second = index("needle", 10)

self.assertEqual(first, second)
index.rebuild([Document("replacement", "different token", {})])
self.assertEqual(index("needle", 10), [])

def test_ranks_overlap_deterministically_and_respects_limit(self) -> None:
index = LocalDocumentIndex(min_chunk_chars=1)
index.rebuild(
[
Document("z", "red blue", {}),
Document("a", "red blue green", {}),
Document("m", "red only", {}),
]
)

hits = index("RED, blue green", 2)

self.assertEqual([hit.metadata["document_id"] for hit in hits], ["a", "z"])
self.assertEqual(index("absent", 3), [])
self.assertEqual(index("red", 0), [])

def test_rejects_duplicate_caller_ids(self) -> None:
index = LocalDocumentIndex(min_chunk_chars=1)

with self.assertRaisesRegex(ValueError, "duplicate document id"):
index.rebuild([Document("same", "one", {}), Document("same", "two", {})])


if __name__ == "__main__":
unittest.main()
Loading