use std::collections::HashSet;
use crate::{
Result,
codec::EncodedVector,
index::{DocumentTokens, Index, build_inverted_file},
};
#[derive(Debug, Clone, Copy)]
pub struct IndexUpdate<'a> {
pub deletions: &'a [u64],
pub upserts: &'a [DocumentTokens],
}
pub fn apply_update(index: Index, update: IndexUpdate<'_>) -> Result<Index> {
let params = index.params;
for doc in update.upserts {
assert!(
doc.tokens.len() == doc.n_tokens * params.dim,
"apply_update: doc {} declared {} tokens but carries {} f32s (dim={})",
doc.doc_id,
doc.n_tokens,
doc.tokens.len(),
params.dim,
);
}
let mut seen: HashSet<u64> = HashSet::with_capacity(update.upserts.len());
for doc in update.upserts {
assert!(
seen.insert(doc.doc_id),
"apply_update: duplicate doc_id {} in upserts",
doc.doc_id,
);
}
let mut to_remove: HashSet<u64> =
update.deletions.iter().copied().collect();
for doc in update.upserts {
to_remove.insert(doc.doc_id);
}
let Index {
codec,
doc_ids,
doc_tokens,
..
} = index;
let mut new_doc_ids: Vec<u64> = Vec::with_capacity(doc_ids.len());
let mut new_doc_tokens: Vec<Vec<EncodedVector>> =
Vec::with_capacity(doc_tokens.len());
for (id, tokens) in doc_ids.into_iter().zip(doc_tokens) {
if to_remove.contains(&id) {
continue;
}
new_doc_ids.push(id);
new_doc_tokens.push(tokens);
}
let total_upsert_tokens: usize =
update.upserts.iter().map(|d| d.n_tokens).sum();
if total_upsert_tokens > 0 {
let mut pool: Vec<f32> =
Vec::with_capacity(total_upsert_tokens * params.dim);
for doc in update.upserts {
pool.extend_from_slice(&doc.tokens);
}
let (all_centroid_ids, all_codes) = codec.batch_encode_tokens(&pool)?;
let packed_per_token = codec.packed_bytes();
let mut offset = 0usize;
for doc in update.upserts {
let n = doc.n_tokens;
let cids = &all_centroid_ids[offset..offset + n];
let codes_slice = &all_codes
[offset * packed_per_token..(offset + n) * packed_per_token];
let encoded: Vec<EncodedVector> = (0..n)
.map(|i| EncodedVector {
centroid_id: cids[i],
codes: codes_slice
[i * packed_per_token..(i + 1) * packed_per_token]
.to_vec(),
})
.collect();
new_doc_ids.push(doc.doc_id);
new_doc_tokens.push(encoded);
offset += n;
}
} else {
for doc in update.upserts {
new_doc_ids.push(doc.doc_id);
new_doc_tokens.push(Vec::new());
}
}
let ivf = build_inverted_file(&new_doc_tokens, params.k_centroids);
Ok(Index {
params,
codec,
doc_ids: new_doc_ids,
doc_tokens: new_doc_tokens,
ivf,
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::index::{IndexParams, build_index};
fn seed_corpus() -> Vec<DocumentTokens> {
vec![
DocumentTokens {
doc_id: 1,
tokens: vec![0.0, 0.0, 0.1, 0.2, -0.1, 0.1],
n_tokens: 3,
},
DocumentTokens {
doc_id: 2,
tokens: vec![10.0, 10.0, 10.2, 9.9, 9.8, 10.1],
n_tokens: 3,
},
DocumentTokens {
doc_id: 3,
tokens: vec![0.3, -0.2, 9.7, 10.2],
n_tokens: 2,
},
]
}
fn seed_params() -> IndexParams {
IndexParams {
dim: 2,
nbits: 2,
k_centroids: 2,
max_kmeans_iters: 50,
}
}
fn seed_index() -> Index {
build_index(&seed_corpus(), seed_params()).unwrap()
}
fn assert_ivf_covers_every_doc_centroid_pair(index: &Index) {
for (doc_idx, encoded) in index.doc_tokens.iter().enumerate() {
for ev in encoded {
let list = index.ivf.docs_for_centroid(ev.centroid_id as usize);
assert!(
list.contains(&(doc_idx as u32)),
"missing posting for doc_idx={doc_idx} centroid={}",
ev.centroid_id,
);
}
}
for c in 0..index.ivf.num_centroids() {
let postings = index.ivf.docs_for_centroid(c);
let mut sorted = postings.to_vec();
sorted.sort_unstable();
sorted.dedup();
assert_eq!(
sorted.len(),
postings.len(),
"duplicate docs in centroid {c} postings",
);
}
}
#[test]
fn apply_update_with_empty_mutations_preserves_every_document() {
let index = seed_index();
let before_ids = index.doc_ids.clone();
let before_tokens = index.doc_tokens.clone();
let updated = apply_update(
index,
IndexUpdate {
deletions: &[],
upserts: &[],
},
)
.unwrap();
assert_eq!(updated.doc_ids, before_ids);
assert_eq!(updated.doc_tokens, before_tokens);
assert_ivf_covers_every_doc_centroid_pair(&updated);
}
#[test]
fn apply_update_removes_the_listed_deletions() {
let index = seed_index();
let original_tokens = index.num_tokens();
let removed_tokens = index
.doc_tokens
.iter()
.zip(&index.doc_ids)
.find(|(_, id)| **id == 2)
.map(|(t, _)| t.len())
.unwrap();
let updated = apply_update(
index,
IndexUpdate {
deletions: &[2],
upserts: &[],
},
)
.unwrap();
assert_eq!(updated.doc_ids, vec![1, 3]);
assert_eq!(updated.doc_tokens.len(), 2);
assert_eq!(updated.num_tokens(), original_tokens - removed_tokens);
assert_ivf_covers_every_doc_centroid_pair(&updated);
}
#[test]
fn apply_update_appends_upsert_of_a_new_doc_id() {
let index = seed_index();
let new_doc = DocumentTokens {
doc_id: 99,
tokens: vec![0.05, -0.05, 0.1, 0.0],
n_tokens: 2,
};
let updated = apply_update(
index,
IndexUpdate {
deletions: &[],
upserts: std::slice::from_ref(&new_doc),
},
)
.unwrap();
assert_eq!(updated.doc_ids, vec![1, 2, 3, 99]);
let encoded = updated.doc_tokens.last().unwrap();
assert_eq!(encoded.len(), 2);
assert_ivf_covers_every_doc_centroid_pair(&updated);
}
#[test]
fn apply_update_replaces_an_existing_doc_when_upserted() {
let index = seed_index();
let replacement = DocumentTokens {
doc_id: 1,
tokens: vec![10.0, 10.1, 9.9, 10.0, 10.1, 9.8, 10.2, 10.0],
n_tokens: 4,
};
let old_encoded = index
.doc_tokens
.iter()
.zip(&index.doc_ids)
.find(|(_, id)| **id == 1)
.map(|(t, _)| t.clone())
.unwrap();
let updated = apply_update(
index,
IndexUpdate {
deletions: &[],
upserts: std::slice::from_ref(&replacement),
},
)
.unwrap();
assert_eq!(updated.doc_ids.len(), 3);
assert_eq!(updated.position_of(1), Some(2));
let new_encoded = &updated.doc_tokens[2];
assert_eq!(new_encoded.len(), 4);
assert_ne!(
*new_encoded, old_encoded,
"upsert must replace the old encoded tokens",
);
assert_ivf_covers_every_doc_centroid_pair(&updated);
}
#[test]
fn apply_update_keeps_the_codec_bit_for_bit() {
let index = seed_index();
let before = index.codec.clone();
let upsert = DocumentTokens {
doc_id: 4,
tokens: vec![5.0, 5.0],
n_tokens: 1,
};
let updated = apply_update(
index,
IndexUpdate {
deletions: &[2],
upserts: std::slice::from_ref(&upsert),
},
)
.unwrap();
assert_eq!(updated.codec.centroids, before.centroids);
assert_eq!(updated.codec.bucket_cutoffs, before.bucket_cutoffs);
assert_eq!(updated.codec.bucket_weights, before.bucket_weights);
assert_eq!(updated.codec.nbits, before.nbits);
assert_eq!(updated.codec.dim, before.dim);
}
#[test]
fn apply_update_preserves_surviving_documents_verbatim() {
let index = seed_index();
let keep_ids: Vec<u64> = index
.doc_ids
.iter()
.copied()
.filter(|id| *id != 2)
.collect();
let keep_tokens: Vec<Vec<EncodedVector>> = index
.doc_ids
.iter()
.zip(&index.doc_tokens)
.filter(|(id, _)| **id != 2)
.map(|(_, t)| t.clone())
.collect();
let updated = apply_update(
index,
IndexUpdate {
deletions: &[2],
upserts: &[],
},
)
.unwrap();
assert_eq!(updated.doc_ids, keep_ids);
assert_eq!(updated.doc_tokens, keep_tokens);
}
#[test]
fn apply_update_handles_upsert_of_an_empty_document() {
let index = seed_index();
let empty = DocumentTokens {
doc_id: 77,
tokens: vec![],
n_tokens: 0,
};
let updated = apply_update(
index,
IndexUpdate {
deletions: &[],
upserts: std::slice::from_ref(&empty),
},
)
.unwrap();
assert!(updated.doc_ids.contains(&77));
assert_eq!(
updated.doc_tokens[updated.position_of(77).unwrap()].len(),
0
);
assert_ivf_covers_every_doc_centroid_pair(&updated);
}
#[test]
#[should_panic(expected = "carries")]
fn apply_update_panics_on_dim_mismatch_in_upsert() {
let index = seed_index();
let bad = DocumentTokens {
doc_id: 5,
tokens: vec![1.0, 2.0, 3.0], n_tokens: 2,
};
let _ = apply_update(
index,
IndexUpdate {
deletions: &[],
upserts: std::slice::from_ref(&bad),
},
)
.unwrap();
}
#[test]
#[should_panic(expected = "duplicate doc_id")]
fn apply_update_panics_on_duplicate_upsert_doc_ids() {
let index = seed_index();
let doc_a = DocumentTokens {
doc_id: 1,
tokens: vec![0.0, 0.0],
n_tokens: 1,
};
let doc_b = DocumentTokens {
doc_id: 1,
tokens: vec![1.0, 1.0],
n_tokens: 1,
};
let _ = apply_update(
index,
IndexUpdate {
deletions: &[],
upserts: &[doc_a, doc_b],
},
)
.unwrap();
}
}