use std::{
fs::File,
io::{BufReader, BufWriter, Read, Write},
path::Path,
};
use crate::{
PlaidError,
Result,
codec::{EncodedVector, ResidualCodec, packed_bytes_per_vector},
index::{Index, IndexParams, build_inverted_file},
};
const MAGIC: &[u8; 8] = b"PLAIDIDX";
const FORMAT_VERSION: u32 = 2;
pub fn save(index: &Index, path: &Path) -> Result<()> {
let file = File::create(path)?;
let mut writer = BufWriter::new(file);
write_index(index, &mut writer)?;
writer.flush()?;
Ok(())
}
pub fn load(path: &Path) -> Result<Index> {
let file = File::open(path)?;
let mut reader = BufReader::new(file);
read_index(&mut reader)
}
fn write_index<W: Write>(index: &Index, w: &mut W) -> Result<()> {
w.write_all(MAGIC)?;
write_u32(w, FORMAT_VERSION)?;
let params = &index.params;
write_u32(w, params.dim as u32)?;
write_u32(w, params.nbits)?;
write_u32(w, params.k_centroids as u32)?;
write_u32(w, params.max_kmeans_iters as u32)?;
write_u64(w, index.doc_ids.len() as u64)?;
write_f32_slice(w, &index.codec.centroids)?;
write_f32_slice(w, &index.codec.bucket_cutoffs)?;
write_f32_slice(w, &index.codec.bucket_weights)?;
write_u64_slice(w, &index.doc_ids)?;
let token_counts: Vec<u32> = index
.doc_tokens
.iter()
.map(|tokens| tokens.len() as u32)
.collect();
write_u32_slice(w, &token_counts)?;
let packed_bytes = packed_bytes_per_vector(params.dim, params.nbits);
for encoded_doc in &index.doc_tokens {
for ev in encoded_doc {
write_u32(w, ev.centroid_id)?;
if ev.codes.len() != packed_bytes {
return Err(PlaidError::InvalidIndex(format!(
"encoded token has {} packed bytes but codec expects {packed_bytes}",
ev.codes.len(),
)));
}
w.write_all(&ev.codes)?;
}
}
Ok(())
}
fn read_index<R: Read>(r: &mut R) -> Result<Index> {
let mut magic = [0u8; 8];
r.read_exact(&mut magic)?;
if &magic != MAGIC {
return Err(PlaidError::InvalidIndex(
"not a docbert-plaid index (magic bytes mismatch)".into(),
));
}
let version = read_u32(r)?;
if version != FORMAT_VERSION {
return Err(PlaidError::InvalidIndex(format!(
"unsupported plaid index version {version}, expected {FORMAT_VERSION}",
)));
}
let dim = read_u32(r)? as usize;
let nbits = read_u32(r)?;
let k_centroids = read_u32(r)? as usize;
let max_kmeans_iters = read_u32(r)? as usize;
let n_documents = read_u64(r)? as usize;
if dim == 0 || k_centroids == 0 || !matches!(nbits, 1 | 2 | 4 | 8) {
return Err(PlaidError::InvalidIndex(
"plaid index header has invalid dim/k_centroids/nbits".into(),
));
}
let params = IndexParams {
dim,
nbits,
k_centroids,
max_kmeans_iters,
};
let centroids = read_f32_vec(r, k_centroids * dim)?;
let num_buckets = 1usize << nbits;
let bucket_cutoffs = read_f32_vec(r, num_buckets - 1)?;
let bucket_weights = read_f32_vec(r, num_buckets)?;
let doc_ids = read_u64_vec(r, n_documents)?;
let token_counts = read_u32_vec(r, n_documents)?;
let packed_bytes = packed_bytes_per_vector(dim, nbits);
let mut doc_tokens: Vec<Vec<EncodedVector>> =
Vec::with_capacity(n_documents);
for count in token_counts.iter() {
let mut encoded_doc = Vec::with_capacity(*count as usize);
for _ in 0..*count {
let centroid_id = read_u32(r)?;
if (centroid_id as usize) >= k_centroids {
return Err(PlaidError::InvalidIndex(format!(
"centroid_id {centroid_id} out of range 0..{k_centroids}",
)));
}
let mut codes = vec![0u8; packed_bytes];
r.read_exact(&mut codes)?;
encoded_doc.push(EncodedVector { centroid_id, codes });
}
doc_tokens.push(encoded_doc);
}
let ivf = build_inverted_file(&doc_tokens, k_centroids);
let codec = ResidualCodec {
nbits,
dim,
centroids,
bucket_cutoffs,
bucket_weights,
};
codec.validate()?;
Ok(Index {
params,
codec,
doc_ids,
doc_tokens,
ivf,
})
}
fn write_u32<W: Write>(w: &mut W, v: u32) -> Result<()> {
w.write_all(&v.to_le_bytes())?;
Ok(())
}
fn write_u64<W: Write>(w: &mut W, v: u64) -> Result<()> {
w.write_all(&v.to_le_bytes())?;
Ok(())
}
fn write_f32_slice<W: Write>(w: &mut W, slice: &[f32]) -> Result<()> {
w.write_all(bytemuck::cast_slice(slice))?;
Ok(())
}
fn write_u32_slice<W: Write>(w: &mut W, slice: &[u32]) -> Result<()> {
w.write_all(bytemuck::cast_slice(slice))?;
Ok(())
}
fn write_u64_slice<W: Write>(w: &mut W, slice: &[u64]) -> Result<()> {
w.write_all(bytemuck::cast_slice(slice))?;
Ok(())
}
fn read_u32<R: Read>(r: &mut R) -> Result<u32> {
let mut buf = [0u8; 4];
r.read_exact(&mut buf)?;
Ok(u32::from_le_bytes(buf))
}
fn read_u64<R: Read>(r: &mut R) -> Result<u64> {
let mut buf = [0u8; 8];
r.read_exact(&mut buf)?;
Ok(u64::from_le_bytes(buf))
}
fn read_f32_vec<R: Read>(r: &mut R, n: usize) -> Result<Vec<f32>> {
let mut out = vec![0.0f32; n];
r.read_exact(bytemuck::cast_slice_mut(&mut out))?;
Ok(out)
}
fn read_u32_vec<R: Read>(r: &mut R, n: usize) -> Result<Vec<u32>> {
let mut out = vec![0u32; n];
r.read_exact(bytemuck::cast_slice_mut(&mut out))?;
Ok(out)
}
fn read_u64_vec<R: Read>(r: &mut R, n: usize) -> Result<Vec<u64>> {
let mut out = vec![0u64; n];
r.read_exact(bytemuck::cast_slice_mut(&mut out))?;
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::index::{DocumentTokens, build_index};
fn small_corpus() -> Vec<DocumentTokens> {
vec![
DocumentTokens {
doc_id: 10,
tokens: vec![0.0, 0.0, 0.1, 0.2, -0.1, 0.1],
n_tokens: 3,
},
DocumentTokens {
doc_id: 20,
tokens: vec![10.0, 10.0, 10.2, 9.9, 9.8, 10.1],
n_tokens: 3,
},
DocumentTokens {
doc_id: 30,
tokens: vec![0.3, -0.2, 9.7, 10.2],
n_tokens: 2,
},
]
}
fn default_params() -> IndexParams {
IndexParams {
dim: 2,
nbits: 2,
k_centroids: 2,
max_kmeans_iters: 50,
}
}
#[test]
fn round_trip_preserves_codec_parameters_and_doc_ids() {
let tmp = tempfile::tempdir().unwrap();
let path = tmp.path().join("index.plaid");
let index = build_index(&small_corpus(), default_params()).unwrap();
save(&index, &path).unwrap();
let loaded = load(&path).unwrap();
assert_eq!(loaded.params.dim, index.params.dim);
assert_eq!(loaded.params.nbits, index.params.nbits);
assert_eq!(loaded.params.k_centroids, index.params.k_centroids);
assert_eq!(
loaded.params.max_kmeans_iters,
index.params.max_kmeans_iters
);
assert_eq!(loaded.doc_ids, index.doc_ids);
}
#[test]
fn round_trip_preserves_codec_tables_byte_for_byte() {
let tmp = tempfile::tempdir().unwrap();
let path = tmp.path().join("index.plaid");
let index = build_index(&small_corpus(), default_params()).unwrap();
save(&index, &path).unwrap();
let loaded = load(&path).unwrap();
assert_eq!(loaded.codec.centroids, index.codec.centroids);
assert_eq!(loaded.codec.bucket_cutoffs, index.codec.bucket_cutoffs);
assert_eq!(loaded.codec.bucket_weights, index.codec.bucket_weights);
}
#[test]
fn round_trip_preserves_encoded_tokens() {
let tmp = tempfile::tempdir().unwrap();
let path = tmp.path().join("index.plaid");
let index = build_index(&small_corpus(), default_params()).unwrap();
save(&index, &path).unwrap();
let loaded = load(&path).unwrap();
assert_eq!(loaded.doc_tokens.len(), index.doc_tokens.len());
for (a, b) in loaded.doc_tokens.iter().zip(index.doc_tokens.iter()) {
assert_eq!(a, b);
}
}
#[test]
fn round_trip_rebuilds_the_inverted_file() {
let tmp = tempfile::tempdir().unwrap();
let path = tmp.path().join("index.plaid");
let index = build_index(&small_corpus(), default_params()).unwrap();
save(&index, &path).unwrap();
let loaded = load(&path).unwrap();
assert_eq!(loaded.ivf.num_centroids(), index.ivf.num_centroids());
assert_eq!(
loaded.ivf.total_doc_postings(),
index.ivf.total_doc_postings(),
);
for c in 0..index.ivf.num_centroids() {
let want = index.ivf.docs_for_centroid(c);
let got = loaded.ivf.docs_for_centroid(c);
assert_eq!(got, want, "IVF list for centroid {c} differs");
}
}
#[test]
fn round_trip_preserves_search_results_exactly() {
let tmp = tempfile::tempdir().unwrap();
let path = tmp.path().join("index.plaid");
let index = build_index(&small_corpus(), default_params()).unwrap();
save(&index, &path).unwrap();
let loaded = load(&path).unwrap();
let query = [0.05f32, 0.1, 9.9, 10.1];
let params = crate::search::SearchParams {
top_k: 3,
n_probe: 2,
n_candidate_docs: None,
centroid_score_threshold: None,
};
let a = crate::search::search(&index, &query, params).unwrap();
let b = crate::search::search(&loaded, &query, params).unwrap();
assert_eq!(a, b);
}
#[test]
fn round_trip_handles_empty_documents() {
let tmp = tempfile::tempdir().unwrap();
let path = tmp.path().join("index.plaid");
let mut docs = small_corpus();
docs.push(DocumentTokens {
doc_id: 999,
tokens: vec![],
n_tokens: 0,
});
let index = build_index(&docs, default_params()).unwrap();
save(&index, &path).unwrap();
let loaded = load(&path).unwrap();
assert_eq!(loaded.doc_ids, index.doc_ids);
assert_eq!(loaded.doc_tokens.last().unwrap().len(), 0);
}
#[test]
fn load_rejects_files_with_wrong_magic() {
let tmp = tempfile::tempdir().unwrap();
let path = tmp.path().join("bogus.plaid");
std::fs::write(&path, b"NOTPLAID").unwrap();
let err = load(&path).unwrap_err();
assert!(
matches!(err, PlaidError::InvalidIndex(ref m) if m.contains("magic")),
"expected InvalidIndex with magic message, got {err:?}",
);
}
#[test]
fn load_rejects_unknown_format_version() {
let tmp = tempfile::tempdir().unwrap();
let path = tmp.path().join("future.plaid");
let mut buf = Vec::new();
buf.extend_from_slice(MAGIC);
buf.extend_from_slice(&999u32.to_le_bytes());
std::fs::write(&path, &buf).unwrap();
let err = load(&path).unwrap_err();
assert!(
matches!(err, PlaidError::InvalidIndex(ref m) if m.contains("version")),
"expected InvalidIndex with version message, got {err:?}",
);
}
#[test]
fn load_rejects_legacy_unpacked_format_version_one() {
let tmp = tempfile::tempdir().unwrap();
let path = tmp.path().join("legacy.plaid");
let mut buf = Vec::new();
buf.extend_from_slice(MAGIC);
buf.extend_from_slice(&1u32.to_le_bytes());
std::fs::write(&path, &buf).unwrap();
let err = load(&path).unwrap_err();
assert!(
matches!(err, PlaidError::InvalidIndex(ref m) if m.contains("version")),
"expected InvalidIndex with version message, got {err:?}",
);
}
}