1use std::collections::HashSet;
24
25use crate::{
26 Result,
27 codec::EncodedVector,
28 index::{
29 DocumentTokens,
30 Index,
31 InvertedFile,
32 build_inverted_file_from_flat,
33 },
34};
35
36#[derive(Debug, Clone, Copy)]
44pub struct IndexUpdate<'a> {
45 pub deletions: &'a [u64],
47 pub upserts: &'a [DocumentTokens],
51}
52
53pub fn apply_update(index: Index, update: IndexUpdate<'_>) -> Result<Index> {
69 let params = index.params;
70
71 for doc in update.upserts {
72 assert!(
73 doc.tokens.len() == doc.n_tokens * params.dim,
74 "apply_update: doc {} declared {} tokens but carries {} f32s (dim={})",
75 doc.doc_id,
76 doc.n_tokens,
77 doc.tokens.len(),
78 params.dim,
79 );
80 }
81
82 let mut seen: HashSet<u64> = HashSet::with_capacity(update.upserts.len());
85 for doc in update.upserts {
86 assert!(
87 seen.insert(doc.doc_id),
88 "apply_update: duplicate doc_id {} in upserts",
89 doc.doc_id,
90 );
91 }
92
93 let mut to_remove: HashSet<u64> =
97 update.deletions.iter().copied().collect();
98 for doc in update.upserts {
99 to_remove.insert(doc.doc_id);
100 }
101
102 let codec = index.codec.clone();
106 let packed_bytes = codec.packed_bytes();
107 let existing_doc_ids = index.doc_ids.clone();
108 let existing_doc_tokens: Vec<Vec<EncodedVector>> = (0..existing_doc_ids
109 .len())
110 .map(|i| index.doc_tokens_vec(i))
111 .collect();
112 drop(index);
113
114 let mut new_doc_ids: Vec<u64> = Vec::with_capacity(existing_doc_ids.len());
116 let mut new_doc_tokens: Vec<Vec<EncodedVector>> =
117 Vec::with_capacity(existing_doc_tokens.len());
118 for (id, tokens) in existing_doc_ids.into_iter().zip(existing_doc_tokens) {
119 if to_remove.contains(&id) {
120 continue;
121 }
122 new_doc_ids.push(id);
123 new_doc_tokens.push(tokens);
124 }
125 let _ = packed_bytes; let total_upsert_tokens: usize =
132 update.upserts.iter().map(|d| d.n_tokens).sum();
133 if total_upsert_tokens > 0 {
134 let mut pool: Vec<f32> =
135 Vec::with_capacity(total_upsert_tokens * params.dim);
136 for doc in update.upserts {
137 pool.extend_from_slice(&doc.tokens);
138 }
139 let (all_centroid_ids, all_codes) = codec.batch_encode_tokens(&pool)?;
140 let packed_per_token = codec.packed_bytes();
141 let mut offset = 0usize;
142 for doc in update.upserts {
143 let n = doc.n_tokens;
144 let cids = &all_centroid_ids[offset..offset + n];
145 let codes_slice = &all_codes
146 [offset * packed_per_token..(offset + n) * packed_per_token];
147 let encoded: Vec<EncodedVector> = (0..n)
148 .map(|i| EncodedVector {
149 centroid_id: cids[i],
150 codes: codes_slice
151 [i * packed_per_token..(i + 1) * packed_per_token]
152 .to_vec(),
153 })
154 .collect();
155 new_doc_ids.push(doc.doc_id);
156 new_doc_tokens.push(encoded);
157 offset += n;
158 }
159 } else {
160 for doc in update.upserts {
163 new_doc_ids.push(doc.doc_id);
164 new_doc_tokens.push(Vec::new());
165 }
166 }
167
168 let mut new_index = Index::from_encoded_docs(
179 params,
180 codec,
181 new_doc_ids,
182 new_doc_tokens,
183 InvertedFile::default(),
184 );
185 new_index.ivf = build_inverted_file_from_flat(
186 &new_index.doc_centroid_ids,
187 &new_index.doc_offsets,
188 params.k_centroids,
189 );
190 Ok(new_index)
191}
192
193#[cfg(test)]
194mod tests {
195 use super::*;
196 use crate::index::{IndexParams, build_index};
197
198 fn seed_corpus() -> Vec<DocumentTokens> {
199 vec![
203 DocumentTokens {
204 doc_id: 1,
205 tokens: vec![0.0, 0.0, 0.1, 0.2, -0.1, 0.1],
206 n_tokens: 3,
207 },
208 DocumentTokens {
209 doc_id: 2,
210 tokens: vec![10.0, 10.0, 10.2, 9.9, 9.8, 10.1],
211 n_tokens: 3,
212 },
213 DocumentTokens {
214 doc_id: 3,
215 tokens: vec![0.3, -0.2, 9.7, 10.2],
216 n_tokens: 2,
217 },
218 ]
219 }
220
221 fn seed_params() -> IndexParams {
222 IndexParams {
223 dim: 2,
224 nbits: 2,
225 k_centroids: 2,
226 max_kmeans_iters: 50,
227 }
228 }
229
230 fn seed_index() -> Index {
231 build_index(&seed_corpus(), seed_params()).unwrap()
232 }
233
234 fn all_doc_tokens(index: &Index) -> Vec<Vec<EncodedVector>> {
240 (0..index.num_documents())
241 .map(|i| index.doc_tokens_vec(i))
242 .collect()
243 }
244
245 fn assert_ivf_covers_every_doc_centroid_pair(index: &Index) {
246 for doc_idx in 0..index.num_documents() {
247 for &cid in index.doc_centroid_ids(doc_idx) {
248 let list = index.ivf.docs_for_centroid(cid as usize);
249 assert!(
250 list.contains(&(doc_idx as u32)),
251 "missing posting for doc_idx={doc_idx} centroid={cid}",
252 );
253 }
254 }
255 for c in 0..index.ivf.num_centroids() {
257 let postings = index.ivf.docs_for_centroid(c);
258 let mut sorted = postings.to_vec();
259 sorted.sort_unstable();
260 sorted.dedup();
261 assert_eq!(
262 sorted.len(),
263 postings.len(),
264 "duplicate docs in centroid {c} postings",
265 );
266 }
267 }
268
269 #[test]
270 fn apply_update_with_empty_mutations_preserves_every_document() {
271 let index = seed_index();
272 let before_ids = index.doc_ids.clone();
273 let before_tokens = all_doc_tokens(&index);
274
275 let updated = apply_update(
276 index,
277 IndexUpdate {
278 deletions: &[],
279 upserts: &[],
280 },
281 )
282 .unwrap();
283
284 assert_eq!(updated.doc_ids, before_ids);
285 assert_eq!(all_doc_tokens(&updated), before_tokens);
286 assert_ivf_covers_every_doc_centroid_pair(&updated);
287 }
288
289 #[test]
290 fn apply_update_removes_the_listed_deletions() {
291 let index = seed_index();
292 let original_tokens = index.num_tokens();
293 let removed_idx = index.position_of(2).unwrap();
294 let removed_tokens = index.doc_token_count(removed_idx);
295
296 let updated = apply_update(
297 index,
298 IndexUpdate {
299 deletions: &[2],
300 upserts: &[],
301 },
302 )
303 .unwrap();
304
305 assert_eq!(updated.doc_ids, vec![1, 3]);
306 assert_eq!(updated.num_documents(), 2);
307 assert_eq!(updated.num_tokens(), original_tokens - removed_tokens);
308 assert_ivf_covers_every_doc_centroid_pair(&updated);
309 }
310
311 #[test]
312 fn apply_update_appends_upsert_of_a_new_doc_id() {
313 let index = seed_index();
314 let new_doc = DocumentTokens {
315 doc_id: 99,
316 tokens: vec![0.05, -0.05, 0.1, 0.0],
317 n_tokens: 2,
318 };
319
320 let updated = apply_update(
321 index,
322 IndexUpdate {
323 deletions: &[],
324 upserts: std::slice::from_ref(&new_doc),
325 },
326 )
327 .unwrap();
328
329 assert_eq!(updated.doc_ids, vec![1, 2, 3, 99]);
330 let last_idx = updated.num_documents() - 1;
331 assert_eq!(updated.doc_token_count(last_idx), 2);
332 assert_ivf_covers_every_doc_centroid_pair(&updated);
333 }
334
335 #[test]
336 fn apply_update_replaces_an_existing_doc_when_upserted() {
337 let index = seed_index();
338 let replacement = DocumentTokens {
339 doc_id: 1,
340 tokens: vec![10.0, 10.1, 9.9, 10.0, 10.1, 9.8, 10.2, 10.0],
343 n_tokens: 4,
344 };
345 let old_encoded = {
346 let idx = index.position_of(1).unwrap();
347 index.doc_tokens_vec(idx)
348 };
349
350 let updated = apply_update(
351 index,
352 IndexUpdate {
353 deletions: &[],
354 upserts: std::slice::from_ref(&replacement),
355 },
356 )
357 .unwrap();
358
359 assert_eq!(updated.doc_ids.len(), 3);
362 assert_eq!(updated.position_of(1), Some(2));
363
364 let new_encoded = updated.doc_tokens_vec(2);
365 assert_eq!(new_encoded.len(), 4);
366 assert_ne!(
367 new_encoded, old_encoded,
368 "upsert must replace the old encoded tokens",
369 );
370 assert_ivf_covers_every_doc_centroid_pair(&updated);
371 }
372
373 #[test]
374 fn apply_update_keeps_the_codec_bit_for_bit() {
375 let index = seed_index();
376 let before = index.codec.clone();
377 let upsert = DocumentTokens {
378 doc_id: 4,
379 tokens: vec![5.0, 5.0],
380 n_tokens: 1,
381 };
382
383 let updated = apply_update(
384 index,
385 IndexUpdate {
386 deletions: &[2],
387 upserts: std::slice::from_ref(&upsert),
388 },
389 )
390 .unwrap();
391
392 assert_eq!(updated.codec.centroids, before.centroids);
394 assert_eq!(updated.codec.bucket_cutoffs, before.bucket_cutoffs);
395 assert_eq!(updated.codec.bucket_weights, before.bucket_weights);
396 assert_eq!(updated.codec.nbits, before.nbits);
397 assert_eq!(updated.codec.dim, before.dim);
398 }
399
400 #[test]
401 fn apply_update_preserves_surviving_documents_verbatim() {
402 let index = seed_index();
403 let keep_ids: Vec<u64> = index
406 .doc_ids
407 .iter()
408 .copied()
409 .filter(|id| *id != 2)
410 .collect();
411 let keep_tokens: Vec<Vec<EncodedVector>> = index
412 .doc_ids
413 .iter()
414 .enumerate()
415 .filter(|(_, id)| **id != 2)
416 .map(|(i, _)| index.doc_tokens_vec(i))
417 .collect();
418
419 let updated = apply_update(
420 index,
421 IndexUpdate {
422 deletions: &[2],
423 upserts: &[],
424 },
425 )
426 .unwrap();
427
428 assert_eq!(updated.doc_ids, keep_ids);
429 assert_eq!(all_doc_tokens(&updated), keep_tokens);
430 }
431
432 #[test]
433 fn apply_update_handles_upsert_of_an_empty_document() {
434 let index = seed_index();
435 let empty = DocumentTokens {
436 doc_id: 77,
437 tokens: vec![],
438 n_tokens: 0,
439 };
440
441 let updated = apply_update(
442 index,
443 IndexUpdate {
444 deletions: &[],
445 upserts: std::slice::from_ref(&empty),
446 },
447 )
448 .unwrap();
449
450 assert!(updated.doc_ids.contains(&77));
451 assert_eq!(
452 updated.doc_token_count(updated.position_of(77).unwrap()),
453 0
454 );
455 assert_ivf_covers_every_doc_centroid_pair(&updated);
456 }
457
458 #[test]
459 #[should_panic(expected = "carries")]
460 fn apply_update_panics_on_dim_mismatch_in_upsert() {
461 let index = seed_index();
462 let bad = DocumentTokens {
463 doc_id: 5,
464 tokens: vec![1.0, 2.0, 3.0], n_tokens: 2,
466 };
467 let _ = apply_update(
468 index,
469 IndexUpdate {
470 deletions: &[],
471 upserts: std::slice::from_ref(&bad),
472 },
473 )
474 .unwrap();
475 }
476
477 #[test]
478 #[should_panic(expected = "duplicate doc_id")]
479 fn apply_update_panics_on_duplicate_upsert_doc_ids() {
480 let index = seed_index();
481 let doc_a = DocumentTokens {
482 doc_id: 1,
483 tokens: vec![0.0, 0.0],
484 n_tokens: 1,
485 };
486 let doc_b = DocumentTokens {
487 doc_id: 1,
488 tokens: vec![1.0, 1.0],
489 n_tokens: 1,
490 };
491 let _ = apply_update(
492 index,
493 IndexUpdate {
494 deletions: &[],
495 upserts: &[doc_a, doc_b],
496 },
497 )
498 .unwrap();
499 }
500}