1use lc_embeddings::token_level::{TokenEmbedding, TokenLevelEmbeddings};
16use lc_embeddings::EmbeddingError;
17
18#[derive(Debug, Clone)]
20pub struct LateChunkConfig {
21 pub chunk_size: usize,
23 pub chunk_overlap: usize,
25}
26
27impl Default for LateChunkConfig {
28 fn default() -> Self {
29 Self {
30 chunk_size: 1024,
31 chunk_overlap: 128,
32 }
33 }
34}
35
36impl LateChunkConfig {
37 pub fn new() -> Self {
39 Self::default()
40 }
41
42 pub fn with_chunk_size(mut self, chunk_size: usize) -> Self {
44 self.chunk_size = chunk_size;
45 self
46 }
47
48 pub fn with_chunk_overlap(mut self, chunk_overlap: usize) -> Self {
50 self.chunk_overlap = chunk_overlap;
51 self
52 }
53
54 pub fn validate(&self) -> Result<(), EmbeddingError> {
57 if self.chunk_size == 0 {
58 return Err(EmbeddingError::Config(
59 "late chunking: chunk_size must be > 0".to_string(),
60 ));
61 }
62 if self.chunk_overlap >= self.chunk_size {
63 return Err(EmbeddingError::Config(format!(
64 "late chunking: chunk_overlap ({}) must be < chunk_size ({})",
65 self.chunk_overlap, self.chunk_size
66 )));
67 }
68 Ok(())
69 }
70
71 pub fn chunk_ranges(&self, text: &str) -> Vec<(usize, usize)> {
79 let text_len = text.len();
80 let mut ranges = Vec::new();
81 let step = self.chunk_size - self.chunk_overlap;
82 let mut start = 0usize;
83 while start < text_len {
84 let raw_end = (start + self.chunk_size).min(text_len);
85 let mut end = prev_char_boundary(text, raw_end);
87 if end <= start {
88 end = next_char_boundary(text, raw_end).min(text_len);
91 }
92 if end <= start {
93 break;
94 }
95 ranges.push((start, end));
96 if end >= text_len {
97 break;
98 }
99 let next = next_char_boundary(text, (start + step).min(text_len));
100 if next <= start {
101 break;
102 }
103 start = next;
104 }
105 ranges
106 }
107}
108
109fn next_char_boundary(text: &str, pos: usize) -> usize {
111 let mut p = pos.min(text.len());
112 while p < text.len() && !text.is_char_boundary(p) {
113 p += 1;
114 }
115 p
116}
117
118fn prev_char_boundary(text: &str, pos: usize) -> usize {
120 let mut p = pos.min(text.len());
121 while p > 0 && !text.is_char_boundary(p) {
122 p -= 1;
123 }
124 p
125}
126
127#[derive(Debug, Clone, PartialEq)]
130pub struct LateChunk {
131 pub range: (usize, usize),
133 pub text: String,
135 pub vector: Vec<f32>,
137}
138
139pub fn pool_tokens(
148 tokens: &[TokenEmbedding],
149 range_start: usize,
150 range_end: usize,
151) -> Result<Vec<f32>, EmbeddingError> {
152 let dim = tokens.first().map(|t| t.vector.len()).unwrap_or(0);
153 let mut pooled = vec![0.0f32; dim];
154 let mut count = 0usize;
155 for token in tokens {
156 if token.span.intersects(range_start, range_end) {
157 for (p, v) in pooled.iter_mut().zip(token.vector.iter()) {
158 *p += v;
159 }
160 count += 1;
161 }
162 }
163 if count == 0 {
164 return Err(EmbeddingError::EmptyInput);
165 }
166 for p in &mut pooled {
167 *p /= count as f32;
168 }
169 lc_embeddings::l2_normalize(&mut pooled);
170 Ok(pooled)
171}
172
173pub async fn late_chunk<E: TokenLevelEmbeddings>(
182 embedder: &E,
183 text: &str,
184 config: &LateChunkConfig,
185) -> Result<Vec<LateChunk>, EmbeddingError> {
186 config.validate()?;
187 let tokens = embedder.embed_tokens(text).await?;
188 let mut chunks = Vec::new();
189 for (start, end) in config.chunk_ranges(text) {
190 let vector = pool_tokens(&tokens, start, end)?;
191 chunks.push(LateChunk {
192 range: (start, end),
193 text: text[start..end].to_string(),
194 vector,
195 });
196 }
197 Ok(chunks)
198}
199
200#[cfg(test)]
201mod tests {
202 use super::*;
203 use lc_embeddings::token_level::{TokenEmbedding, TokenSpan};
204
205 fn token_embeddings(text: &str) -> Vec<TokenEmbedding> {
208 let mut out = Vec::new();
209 let mut cursor = 0usize;
210 for word in text.split_whitespace() {
211 let start = text[cursor..]
212 .find(word)
213 .map(|p| cursor + p)
214 .unwrap_or(cursor);
215 let end = start + word.len();
216 cursor = end;
217 out.push(TokenEmbedding {
218 span: TokenSpan::new(start, end),
219 vector: vec![word.bytes().map(|b| b as f32).sum::<f32>(), 1.0],
220 });
221 }
222 out
223 }
224
225 #[test]
226 fn config_validates_overlap() {
227 assert!(LateChunkConfig::new().validate().is_ok());
228 let bad = LateChunkConfig {
229 chunk_size: 10,
230 chunk_overlap: 10,
231 };
232 assert!(bad.validate().is_err());
233 let bad = LateChunkConfig {
234 chunk_size: 0,
235 chunk_overlap: 0,
236 };
237 assert!(bad.validate().is_err());
238 }
239
240 #[test]
242 fn chunk_ranges_cover_document() {
243 let config = LateChunkConfig {
244 chunk_size: 10,
245 chunk_overlap: 2,
246 };
247 let text_ascii = "x".repeat(25);
249 let ranges = config.chunk_ranges(&text_ascii);
250 assert_eq!(ranges, vec![(0, 10), (8, 18), (16, 25)]);
252 }
253
254 #[test]
255 fn chunk_ranges_shorter_than_chunk_size() {
256 let config = LateChunkConfig {
257 chunk_size: 10,
258 chunk_overlap: 2,
259 };
260 assert_eq!(config.chunk_ranges("xxxxx"), vec![(0, 5)]);
261 }
262
263 #[test]
266 fn chunk_ranges_snap_to_char_boundaries() {
267 let config = LateChunkConfig {
268 chunk_size: 10,
269 chunk_overlap: 2,
270 };
271 let text = "你好世界天地".to_string(); let ranges = config.chunk_ranges(&text);
274 assert!(!ranges.is_empty());
275 for (start, end) in &ranges {
276 assert!(
277 text.is_char_boundary(*start),
278 "start {start} not a boundary"
279 );
280 assert!(text.is_char_boundary(*end), "end {end} not a boundary");
281 let _ = &text[*start..*end];
283 }
284 assert_eq!(ranges.last().unwrap().1, 18);
286 }
287
288 #[test]
290 fn pool_tokens_averages_and_normalizes() {
291 let tokens = token_embeddings("alpha beta");
292 let pooled = pool_tokens(&tokens, 0, 10).unwrap();
294 let mean0: f32 = (519.0 + 412.0) / 2.0;
295 let norm = (mean0 * mean0 + 1.0).sqrt();
296 assert!((pooled[0] - mean0 / norm).abs() < 1e-5);
297 assert!((pooled[1] - 1.0 / norm).abs() < 1e-5);
298 let norm_sq: f32 = pooled.iter().map(|v| v * v).sum();
299 assert!((norm_sq - 1.0).abs() < 1e-5, "L2-normalized");
300 }
301
302 #[test]
305 fn pool_tokens_boundary_leaks() {
306 let tokens = token_embeddings("abcdef");
307 let spanning = vec![TokenEmbedding {
310 span: TokenSpan::new(2, 5),
311 vector: vec![1.0, 0.0],
312 }];
313 assert!(pool_tokens(&spanning, 0, 3).is_ok());
314 assert!(pool_tokens(&spanning, 3, 6).is_ok());
315 let _ = tokens; }
317
318 #[test]
320 fn pool_tokens_empty_range_errors() {
321 let tokens = token_embeddings("alpha");
322 let err = pool_tokens(&tokens, 100, 200).unwrap_err();
323 assert!(matches!(err, EmbeddingError::EmptyInput));
324 }
325
326 #[tokio::test]
329 async fn late_chunk_end_to_end() {
330 struct MockTokenEmbeddings;
331
332 impl lc_embeddings::token_level::TokenLevelEmbeddings for MockTokenEmbeddings {
333 async fn embed_tokens(
334 &self,
335 text: &str,
336 ) -> Result<Vec<TokenEmbedding>, EmbeddingError> {
337 Ok(token_embeddings(text))
338 }
339 }
340
341 let text = "alpha beta gamma delta epsilon";
342 let config = LateChunkConfig {
343 chunk_size: 16,
344 chunk_overlap: 0,
345 };
346 let chunks = late_chunk(&MockTokenEmbeddings, text, &config)
347 .await
348 .unwrap();
349 assert_eq!(chunks.len(), 2);
350 assert_eq!(chunks[0].text, "alpha beta gamma");
351 assert_eq!(chunks[1].text, " delta epsilon");
352 for chunk in &chunks {
354 let norm_sq: f32 = chunk.vector.iter().map(|v| v * v).sum();
355 assert!((norm_sq - 1.0).abs() < 1e-4);
356 }
357 let delta_sum = b"delta".iter().map(|b| *b as f32).sum::<f32>();
359 let eps_sum: f32 = b"epsilon".iter().map(|b| *b as f32).sum();
360 let mean = (delta_sum + eps_sum) / 2.0;
361 let norm = (mean * mean + 1.0f32).sqrt();
362 assert!((chunks[1].vector[0] - mean / norm).abs() < 1e-4);
363 }
364
365 #[tokio::test]
367 async fn late_chunk_rejects_invalid_config() {
368 struct MockTokenEmbeddings;
369
370 impl lc_embeddings::token_level::TokenLevelEmbeddings for MockTokenEmbeddings {
371 async fn embed_tokens(
372 &self,
373 _text: &str,
374 ) -> Result<Vec<TokenEmbedding>, EmbeddingError> {
375 Ok(Vec::new())
376 }
377 }
378 let config = LateChunkConfig {
379 chunk_size: 8,
380 chunk_overlap: 8,
381 };
382 let err = late_chunk(&MockTokenEmbeddings, "text", &config)
383 .await
384 .unwrap_err();
385 assert!(matches!(err, EmbeddingError::Config(_)));
386 }
387}