This commit is contained in:
Liang Xu 2026-09-13 05:02:37 +08:00 committed by GitHub
commit 40ca5b10a5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 59 additions and 1 deletions

View file

@ -23,6 +23,15 @@ class RecursiveCharacterTextSplitter:
chunk_overlap: int = DEFAULT_CHUNK_OVERLAP,
separators: list[str] | None = None,
):
if chunk_size <= 0:
raise ValueError(f"chunk_size must be greater than 0, got {chunk_size}")
if chunk_overlap < 0:
raise ValueError(f"chunk_overlap must be non-negative, got {chunk_overlap}")
if chunk_overlap >= chunk_size:
raise ValueError(
f"Got a larger or equal chunk_overlap ({chunk_overlap}) than chunk_size ({chunk_size}), "
f"chunk_overlap must be smaller than chunk_size"
)
self.chunk_size = chunk_size
self.chunk_overlap = chunk_overlap
self.separators = separators or ["\n\n", "\n", " ", ""]
@ -123,12 +132,15 @@ class RecursiveCharacterTextSplitter:
"""Force split text by chunk_size when no separator works."""
chunks: Final[list[str]] = []
start = 0
step: Final = max(self.chunk_size - self.chunk_overlap, 1)
while start < len(text):
end = start + self.chunk_size
chunk = text[start:end].strip()
if chunk:
chunks.append(chunk)
start = end - self.chunk_overlap if end < len(text) else len(text)
if end >= len(text):
break
start += step
return chunks

View file

@ -0,0 +1,46 @@
import pytest
from litellm.rag.text_splitters.recursive_character_text_splitter import (
RecursiveCharacterTextSplitter,
)
def test_invalid_chunk_overlap_greater_than_or_equal_chunk_size():
"""chunk_overlap >= chunk_size must raise ValueError to prevent infinite loops."""
with pytest.raises(ValueError, match="chunk_overlap"):
RecursiveCharacterTextSplitter(chunk_size=500, chunk_overlap=500)
with pytest.raises(ValueError, match="chunk_overlap"):
RecursiveCharacterTextSplitter(chunk_size=500, chunk_overlap=1000)
def test_invalid_chunk_size_non_positive():
"""chunk_size <= 0 must raise ValueError."""
with pytest.raises(ValueError, match="chunk_size"):
RecursiveCharacterTextSplitter(chunk_size=0, chunk_overlap=0)
with pytest.raises(ValueError, match="chunk_size"):
RecursiveCharacterTextSplitter(chunk_size=-10, chunk_overlap=0)
def test_invalid_chunk_overlap_negative():
"""chunk_overlap < 0 must raise ValueError."""
with pytest.raises(ValueError, match="chunk_overlap"):
RecursiveCharacterTextSplitter(chunk_size=100, chunk_overlap=-1)
def test_split_text_basic():
splitter = RecursiveCharacterTextSplitter(chunk_size=50, chunk_overlap=10)
text = "This is a test paragraph.\n\nThis is another paragraph that is longer and will need splitting."
chunks = splitter.split_text(text)
assert len(chunks) > 1
for chunk in chunks:
assert len(chunk) <= 50
def test_force_split_no_matching_separators():
"""Verify that text with no matching separators splits cleanly without hanging."""
splitter = RecursiveCharacterTextSplitter(chunk_size=20, chunk_overlap=5, separators=["\n\n"])
text = "abcdefghijklmnopqrstuvwxyz0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZ"
chunks = splitter.split_text(text)
assert len(chunks) > 1
assert "".join(chunks) != ""