From 119382e869fbebab2403e037a77a77dadb511163 Mon Sep 17 00:00:00 2001 From: prashantshukla01 Date: Sat, 26 Sep 2026 12:24:00 +0530 Subject: [PATCH 1/2] feat: add TokenChunker for token-level chunking via tiktoken (closes #44) --- CHANGELOG.md | 3 + pyproject.toml | 2 + ragframework/document/__init__.py | 8 +- ragframework/document/chunkers.py | 73 +++++++++++++++++++ tests/test_document/test_chunkers.py | 105 +++++++++++++++++++++++++++ 5 files changed, 190 insertions(+), 1 deletion(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index e6be596..5c92fe9 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,9 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Added +- `TokenChunker` for token-level chunking using tiktoken with a new `[tokens]` optional extra (closes #44) + ## [0.3.0] - 2026-09-25 ### Added diff --git a/pyproject.toml b/pyproject.toml index 9ce582a..062baa9 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -36,6 +36,7 @@ anthropic = ["anthropic>=0.25"] huggingface = ["sentence-transformers>=2.7"] chromadb = ["chromadb>=0.5"] faiss = ["faiss-cpu>=1.8"] +tokens = ["tiktoken>=0.7"] all = [ "ragframework[pdf]", "ragframework[docx]", @@ -44,6 +45,7 @@ all = [ "ragframework[huggingface]", "ragframework[chromadb]", "ragframework[faiss]", + "ragframework[tokens]", ] dev = [ "pytest>=8.0", diff --git a/ragframework/document/__init__.py b/ragframework/document/__init__.py index 930840e..7f5977d 100644 --- a/ragframework/document/__init__.py +++ b/ragframework/document/__init__.py @@ -1,6 +1,11 @@ """Document loading and chunking utilities.""" -from ragframework.document.chunkers import FixedSizeChunker, RecursiveChunker, SentenceChunker +from ragframework.document.chunkers import ( + FixedSizeChunker, + RecursiveChunker, + SentenceChunker, + TokenChunker, +) from .html import HTMLLoader from .loaders import MarkdownLoader, PDFLoader, TextFileLoader @@ -16,4 +21,5 @@ "FixedSizeChunker", "RecursiveChunker", "SentenceChunker", + "TokenChunker", ] diff --git a/ragframework/document/chunkers.py b/ragframework/document/chunkers.py index 07e2205..b9c5d0f 100644 --- a/ragframework/document/chunkers.py +++ b/ragframework/document/chunkers.py @@ -343,3 +343,76 @@ def _merge_pieces(self, pieces: list[str], next_separators: list[str]) -> list[s final_chunks.append("".join(current_chunk_pieces)) return final_chunks + + +class TokenChunker(TextChunker): + """Split text into chunks measured in tokens using tiktoken. + Args: + chunk_tokens: Maximum number of tokens per chunk. + overlap_tokens: Number of overlapping tokens between adjacent chunks. + encoding_name: Name of the tiktoken encoding to use. + Raises: + ImportError: If the ``tokens`` optional dependency is not installed. + """ + + def __init__( + self, + chunk_tokens: int = 256, + overlap_tokens: int = 32, + encoding_name: str = "cl100k_base", + ) -> None: + # --- Guarded import (same pattern as HuggingFaceEmbedder) --- + try: + import tiktoken + except ImportError as exc: + raise ImportError( + "Token chunking requires 'ragframework[tokens]'. " + "Install it with: pip install ragframework[tokens]" + ) from exc + # --- Validation (same rules as RecursiveChunker) --- + if chunk_tokens <= 0: + raise ValueError("chunk_tokens must be positive") + if overlap_tokens < 0: + raise ValueError("overlap_tokens must be non-negative") + + if overlap_tokens >= chunk_tokens: + raise ValueError("overlap_tokens must be less than chunk_tokens") + + self.chunk_tokens = chunk_tokens + self.overlap_tokens = overlap_tokens + self.encoding_name = encoding_name + self._encoding = tiktoken.get_encoding(encoding_name) + + def chunk(self, document: Document) -> list[Chunk]: + if not document.content: + return [] + + # encode the full text once into token IDs + + all_tokens = self._encoding.encode(document.content) + chunks: list[Chunk] = [] + step = self.chunk_tokens - self.overlap_tokens + index = 0 + chunk_num = 0 + + while index < len(all_tokens): + # slice the token window + window = all_tokens[index : index + self.chunk_tokens] + # decode back to text + chunk_text = self._encoding.decode(window) + token_count = len(window) + + chunks.append( + Chunk( + id=f"{document.id}:{chunk_num}", + content=chunk_text, + metadata={ + **document.metadata, + "chunk_index": chunk_num, + "token_count": token_count, + }, + ) + ) + chunk_num += 1 + index += step + return chunks diff --git a/tests/test_document/test_chunkers.py b/tests/test_document/test_chunkers.py index 617b759..e090ffd 100644 --- a/tests/test_document/test_chunkers.py +++ b/tests/test_document/test_chunkers.py @@ -1,5 +1,9 @@ """Tests for built-in text chunkers.""" +import builtins +import sys +import types + import pytest from ragframework.base import Document @@ -243,3 +247,104 @@ def test_redundant_trailing_chunks(self): ) chunks = chunker.chunk(doc) assert [len(c.content) for c in chunks] == [6, 6, 6] + + +# --------------- Fake tiktoken for CI --------------- + + +class FakeEncoding: + """Trivially encodes each character as one token (its ordinal).""" + + def encode(self, text: str) -> list[int]: + return list(text.encode("utf-8")) + + def decode(self, tokens: list[int]) -> str: + return bytes(tokens).decode("utf-8") + + +@pytest.fixture +def fake_tiktoken(monkeypatch): + """Inject a fake tiktoken module so tests don't need a real download.""" + module = types.ModuleType("tiktoken") + module.get_encoding = lambda name: FakeEncoding() + monkeypatch.setitem(sys.modules, "tiktoken", module) + + +class TestTokenChunker: + def test_basic_chunking(self, fake_tiktoken): + from ragframework.document.chunkers import TokenChunker + + doc = Document(id="d1", content="A" * 100, metadata={}) + chunker = TokenChunker(chunk_tokens=30, overlap_tokens=5) + chunks = chunker.chunk(doc) + + assert len(chunks) > 1 + # No chunk should exceed the token limit + for c in chunks: + assert c.metadata["token_count"] <= 30 + + def test_chunk_ids_follow_convention(self, fake_tiktoken): + from ragframework.document.chunkers import TokenChunker + + doc = Document(id="doc1", content="Hello world this is a test", metadata={"src": "test"}) + chunker = TokenChunker(chunk_tokens=10, overlap_tokens=2) + chunks = chunker.chunk(doc) + + for i, c in enumerate(chunks): + assert c.id == f"doc1:{i}" + assert c.metadata["chunk_index"] == i + assert c.metadata["src"] == "test" + assert "token_count" in c.metadata + + def test_empty_doc_returns_empty(self, fake_tiktoken): + from ragframework.document.chunkers import TokenChunker + + doc = Document(id="x", content="", metadata={}) + chunker = TokenChunker(chunk_tokens=10, overlap_tokens=0) + assert chunker.chunk(doc) == [] + + def test_short_doc_single_chunk(self, fake_tiktoken): + from ragframework.document.chunkers import TokenChunker + + doc = Document(id="x", content="Hi", metadata={}) + chunker = TokenChunker(chunk_tokens=256, overlap_tokens=0) + chunks = chunker.chunk(doc) + assert len(chunks) == 1 + assert chunks[0].content == "Hi" + assert chunks[0].metadata["token_count"] == 2 # 'H' and 'i' in our fake + + def test_validation_errors(self, fake_tiktoken): + from ragframework.document.chunkers import TokenChunker + + with pytest.raises(ValueError, match="chunk_tokens must be positive"): + TokenChunker(chunk_tokens=0) + with pytest.raises(ValueError, match="overlap_tokens must be non-negative"): + TokenChunker(chunk_tokens=10, overlap_tokens=-1) + with pytest.raises(ValueError, match="overlap_tokens must be less than chunk_tokens"): + TokenChunker(chunk_tokens=10, overlap_tokens=10) + + def test_no_chunk_exceeds_token_limit(self, fake_tiktoken): + from ragframework.document.chunkers import TokenChunker + + doc = Document(id="d1", content="abcdefghijklmnopqrstuvwxyz" * 10, metadata={}) + chunker = TokenChunker(chunk_tokens=50, overlap_tokens=10) + chunks = chunker.chunk(doc) + + for c in chunks: + assert c.metadata["token_count"] <= 50 + + def test_missing_tiktoken_gives_helpful_message(self, monkeypatch): + real_import = builtins.__import__ + + def fake_import(name, globals=None, locals=None, fromlist=(), level=0): + if name == "tiktoken": + raise ImportError("No module named 'tiktoken'") + return real_import(name, globals, locals, fromlist, level) + + monkeypatch.delitem(sys.modules, "tiktoken", raising=False) + monkeypatch.setattr(builtins, "__import__", fake_import) + + from ragframework.document.chunkers import TokenChunker + + with pytest.raises(ImportError, match=r"ragframework\[tokens\]"): + TokenChunker() From 5f3ac6597de387a2e0c272b0e2f3517be2b30c0c Mon Sep 17 00:00:00 2001 From: prashantshukla01 Date: Sun, 27 Sep 2026 11:56:19 +0530 Subject: [PATCH 2/2] fix(TokenChunker): preserve multi-byte UTF-8 boundaries and handle special tokens --- ragframework/document/chunkers.py | 54 +++++++++++++++++------- tests/test_document/test_chunkers.py | 63 +++++++++++++++++++++++++--- 2 files changed, 97 insertions(+), 20 deletions(-) diff --git a/ragframework/document/chunkers.py b/ragframework/document/chunkers.py index 2d1fee3..647aab3 100644 --- a/ragframework/document/chunkers.py +++ b/ragframework/document/chunkers.py @@ -402,22 +402,32 @@ def __init__( def chunk(self, document: Document) -> list[Chunk]: if not document.content: return [] - - # encode the full text once into token IDs - - all_tokens = self._encoding.encode(document.content) + # Treat special tokens like <|endoftext|> as ordinary text + all_tokens = self._encoding.encode(document.content, disallowed_special=()) + if not all_tokens: + return [] chunks: list[Chunk] = [] - step = self.chunk_tokens - self.overlap_tokens - index = 0 + start_idx = 0 chunk_num = 0 - - while index < len(all_tokens): - # slice the token window - window = all_tokens[index : index + self.chunk_tokens] - # decode back to text - chunk_text = self._encoding.decode(window) - token_count = len(window) - + total_tokens = len(all_tokens) + while start_idx < total_tokens: + end_idx = min(start_idx + self.chunk_tokens, total_tokens) + # Shrink window backwards if it cuts across a multi-byte UTF-8 character + chunk_text: str | None = None + while end_idx > start_idx: + raw_bytes = self._encoding.decode_bytes(all_tokens[start_idx:end_idx]) + try: + chunk_text = raw_bytes.decode("utf-8") + break + except UnicodeDecodeError: + end_idx -= 1 + # If even 1 token cannot complete a character, it exceeds chunk_tokens + if end_idx == start_idx or chunk_text is None: + raise ValueError( + f"Character at token index {start_idx} cannot fit within " + f"chunk_tokens={self.chunk_tokens}. Increase chunk_tokens." + ) + token_count = end_idx - start_idx chunks.append( Chunk( id=f"{document.id}:{chunk_num}", @@ -430,5 +440,19 @@ def chunk(self, document: Document) -> list[Chunk]: ) ) chunk_num += 1 - index += step + if end_idx >= total_tokens: + break + # Calculate next start_idx respecting overlap and character boundaries + if self.overlap_tokens == 0: + start_idx = end_idx + else: + target_start = max(start_idx + 1, end_idx - self.overlap_tokens) + while target_start < end_idx: + overlap_bytes = self._encoding.decode_bytes(all_tokens[target_start:end_idx]) + try: + overlap_bytes.decode("utf-8") + break + except UnicodeDecodeError: + target_start += 1 + start_idx = target_start return chunks diff --git a/tests/test_document/test_chunkers.py b/tests/test_document/test_chunkers.py index bb274bf..bc0ad00 100644 --- a/tests/test_document/test_chunkers.py +++ b/tests/test_document/test_chunkers.py @@ -12,6 +12,7 @@ RecursiveChunker, SentenceChunker, ) +from typing import Any def test_recursive_chunker_from_config(): @@ -303,13 +304,26 @@ def test_redundant_trailing_chunks(self): class FakeEncoding: - """Trivially encodes each character as one token (its ordinal).""" - - def encode(self, text: str) -> list[int]: + """Emulates tiktoken with byte-level tokens, special token checks, and decode_bytes.""" + + def encode( + self, + text: str, + *, + allowed_special: Any = (), + disallowed_special: Any = "all", + ) -> list[int]: + if disallowed_special and "<|endoftext|>" in text: + raise ValueError( + "Encountered text corresponding to disallowed special token '<|endoftext|>'." + ) return list(text.encode("utf-8")) - def decode(self, tokens: list[int]) -> str: - return bytes(tokens).decode("utf-8") + def decode(self, tokens: list[int], errors: str = "replace") -> str: + return bytes(tokens).decode("utf-8", errors=errors) + + def decode_bytes(self, tokens: list[int]) -> bytes: + return bytes(tokens) @pytest.fixture @@ -398,3 +412,42 @@ def fake_import(name, globals=None, locals=None, fromlist=(), level=0): with pytest.raises(ImportError, match=r"ragframework\[tokens\]"): TokenChunker() + + def test_special_tokens_treated_as_ordinary_text(self, fake_tiktoken): + from ragframework.document.chunkers import TokenChunker + + doc = Document(id="x", content="Hello <|endoftext|> world", metadata={}) + chunker = TokenChunker(chunk_tokens=50, overlap_tokens=0) + chunks = chunker.chunk(doc) + assert len(chunks) == 1 + assert chunks[0].content == "Hello <|endoftext|> world" + + def test_emoji_preserves_complete_characters(self, fake_tiktoken): + from ragframework.document.chunkers import TokenChunker + + doc = Document(id="x", content="A 🙂 B 🚀 C", metadata={}) + chunker = TokenChunker(chunk_tokens=6, overlap_tokens=0) + chunks = chunker.chunk(doc) + assert len(chunks) > 1 + for c in chunks: + assert "\ufffd" not in c.content + assert "".join(c.content for c in chunks) == "A 🙂 B 🚀 C" + + def test_cjk_preserves_complete_characters(self, fake_tiktoken): + from ragframework.document.chunkers import TokenChunker + + doc = Document(id="x", content="你好世界,这是一个测试", metadata={}) + chunker = TokenChunker(chunk_tokens=9, overlap_tokens=0) + chunks = chunker.chunk(doc) + assert len(chunks) > 1 + for c in chunks: + assert "\ufffd" not in c.content + assert "".join(c.content for c in chunks) == "你好世界,这是一个测试" + + def test_character_cannot_fit_raises_value_error(self, fake_tiktoken): + from ragframework.document.chunkers import TokenChunker + + doc = Document(id="x", content="🙂", metadata={}) + chunker = TokenChunker(chunk_tokens=1, overlap_tokens=0) + with pytest.raises(ValueError, match="cannot fit within chunk_tokens"): + chunker.chunk(doc)