mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-13 23:11:40 +00:00
Merge 177ad4ce2a into 559247fa84
This commit is contained in:
commit
40ca5b10a5
2 changed files with 59 additions and 1 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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) != ""
|
||||
Loading…
Add table
Reference in a new issue