lc_shared/
splitter_types.rs1use crate::document_types::Document;
9
10pub trait TextSplitter: Send + Sync {
12 fn split_text(&self, text: &str) -> Vec<String>;
14
15 fn split_document(&self, document: &Document) -> Vec<Document> {
17 let chunks = self.split_text(&document.content);
18 chunks
19 .into_iter()
20 .enumerate()
21 .map(|(i, chunk)| {
22 let mut metadata = document.metadata.clone();
23 metadata.insert("chunk".to_string(), i.to_string());
24
25 Document {
26 content: chunk,
27 metadata,
28 id: None,
29 }
30 })
31 .collect()
32 }
33}
34
35pub struct RecursiveCharacterSplitter {
39 chunk_size: usize,
41
42 chunk_overlap: usize,
44
45 separators: Vec<String>,
47}
48
49impl RecursiveCharacterSplitter {
50 pub fn new(chunk_size: usize, chunk_overlap: usize) -> Self {
52 Self {
53 chunk_size,
54 chunk_overlap,
55 separators: vec![
56 "\n\n".to_string(), "\n".to_string(), "。".to_string(), ".".to_string(), " ".to_string(), "".to_string(), ],
63 }
64 }
65
66 pub fn with_defaults() -> Self {
68 Self::new(1000, 200)
69 }
70
71 pub fn with_separators(mut self, separators: Vec<String>) -> Self {
73 self.separators = separators;
74 self
75 }
76
77 fn split_text_recursive(&self, text: &str, separators: &[String]) -> Vec<String> {
79 let mut chunks = Vec::new();
80
81 if text.is_empty() {
82 return chunks;
83 }
84
85 if text.chars().count() <= self.chunk_size {
87 chunks.push(text.to_string());
88 return chunks;
89 }
90
91 let separator = separators
93 .iter()
94 .find(|s| text.contains(s.as_str()))
95 .cloned()
96 .unwrap_or_default();
97
98 let splits: Vec<String> = if separator.is_empty() {
100 text.chars().map(|c| c.to_string()).collect()
101 } else {
102 text.split(&separator).map(|s| s.to_string()).collect()
103 };
104
105 let mut current_chunk = String::new();
107
108 for split in splits {
109 let split_with_sep = if separator.is_empty() {
110 split.clone()
111 } else if current_chunk.is_empty() {
112 split
113 } else {
114 format!("{}{}", separator, split)
115 };
116
117 if split_with_sep.chars().count() > self.chunk_size {
119 if !current_chunk.is_empty() {
121 chunks.push(current_chunk.clone());
122 current_chunk.clear();
123 }
124
125 let next_separators = if separators.len() > 1 {
127 &separators[1..]
128 } else {
129 &[]
130 };
131
132 let sub_chunks = self.split_text_recursive(&split_with_sep, next_separators);
133 chunks.extend(sub_chunks);
134 } else if current_chunk.chars().count() + split_with_sep.chars().count()
135 > self.chunk_size
136 {
137 chunks.push(current_chunk.clone());
139 current_chunk = split_with_sep;
140 } else {
141 current_chunk.push_str(&split_with_sep);
142 }
143 }
144
145 if !current_chunk.is_empty() {
146 chunks.push(current_chunk);
147 }
148
149 chunks
150 }
151}
152
153impl TextSplitter for RecursiveCharacterSplitter {
154 fn split_text(&self, text: &str) -> Vec<String> {
155 let mut chunks = self.split_text_recursive(text, &self.separators);
156
157 if self.chunk_overlap > 0 && chunks.len() > 1 {
159 let mut overlapped = Vec::new();
160
161 for (i, chunk) in chunks.into_iter().enumerate() {
162 if i == 0 {
163 overlapped.push(chunk);
164 } else {
165 let prev = &overlapped[i - 1];
167 let chars: Vec<char> = prev.chars().collect();
168 let overlap_chars = chars.len().saturating_sub(self.chunk_overlap);
169 let overlap: String = chars[overlap_chars..].iter().collect();
170
171 overlapped.push(format!("{}{}", overlap, chunk));
172 }
173 }
174
175 chunks = overlapped;
176 }
177
178 chunks
179 }
180}
181
182#[allow(dead_code)]
184pub struct CharacterTextSplitter {
185 chunk_size: usize,
187
188 chunk_overlap: usize,
190
191 separator: String,
193}
194
195#[allow(dead_code)]
196impl CharacterTextSplitter {
197 pub fn new(chunk_size: usize, chunk_overlap: usize, separator: &str) -> Self {
199 Self {
200 chunk_size,
201 chunk_overlap,
202 separator: separator.to_string(),
203 }
204 }
205}
206
207impl TextSplitter for CharacterTextSplitter {
208 fn split_text(&self, text: &str) -> Vec<String> {
209 let splits: Vec<&str> = text.split(&self.separator).collect();
210 let mut chunks = Vec::new();
211 let mut current = String::new();
212
213 for split in splits {
214 if current.chars().count() + split.chars().count() + self.separator.chars().count()
215 > self.chunk_size
216 && !current.is_empty()
217 {
218 chunks.push(current.clone());
219
220 if self.chunk_overlap > 0 {
222 let chars: Vec<char> = current.chars().collect();
223 let overlap_start = chars.len().saturating_sub(self.chunk_overlap);
224 current = chars[overlap_start..].iter().collect();
225 } else {
226 current.clear();
227 }
228 }
229
230 if !current.is_empty() {
231 current.push_str(&self.separator);
232 }
233 current.push_str(split);
234 }
235
236 if !current.is_empty() {
237 chunks.push(current);
238 }
239
240 chunks
241 }
242}
243
244#[cfg(test)]
245mod tests {
246 use super::*;
247
248 #[test]
249 fn test_recursive_splitter() {
250 let splitter = RecursiveCharacterSplitter::new(50, 10);
251
252 let text = "This is a sentence. This is another sentence. And a third one.";
253 let chunks = splitter.split_text(text);
254
255 assert!(!chunks.is_empty());
256 for chunk in &chunks {
257 assert!(chunk.len() <= 60); }
259 }
260
261 #[test]
262 fn test_split_document() {
263 let splitter = RecursiveCharacterSplitter::new(100, 20);
264
265 let doc = Document::new("First paragraph.\n\nSecond paragraph.\n\nThird paragraph.")
266 .with_metadata("source", "test");
267
268 let chunks = splitter.split_document(&doc);
269
270 assert!(!chunks.is_empty());
271 for (i, chunk) in chunks.iter().enumerate() {
272 assert!(chunk.metadata.contains_key("chunk"));
273 assert_eq!(chunk.metadata.get("chunk"), Some(&i.to_string()));
274 assert_eq!(chunk.metadata.get("source"), Some(&"test".to_string()));
275 }
276 }
277
278 #[test]
279 fn test_character_splitter() {
280 let splitter = CharacterTextSplitter::new(20, 5, " ");
281
282 let text = "This is a test sentence with multiple words";
283 let chunks = splitter.split_text(text);
284
285 assert!(!chunks.is_empty());
286 }
287
288 #[test]
289 fn test_empty_text() {
290 let splitter = RecursiveCharacterSplitter::new(100, 20);
291 let chunks = splitter.split_text("");
292 assert!(chunks.is_empty());
293 }
294
295 #[test]
296 fn test_small_text() {
297 let splitter = RecursiveCharacterSplitter::new(1000, 200);
298 let text = "Short text";
299 let chunks = splitter.split_text(text);
300 assert_eq!(chunks.len(), 1);
301 assert_eq!(chunks[0], "Short text");
302 }
303}