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