use std::{
fs::File,
io::{BufReader, BufWriter, Read, Write},
path::Path,
};
use crate::{
PlaidError,
Result,
codec::{ResidualCodec, packed_bytes_per_vector},
index::{Index, IndexParams, build_inverted_file_from_flat},
};
const MAGIC: &[u8; 8] = b"PLAIDIDX";
const FORMAT_VERSION: u32 = 3;
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> = (0..index.num_documents())
.map(|i| index.doc_token_count(i) as u32)
.collect();
write_u32_slice(w, &token_counts)?;
let packed_bytes = packed_bytes_per_vector(params.dim, params.nbits);
if index.doc_residual_bytes.len() != index.num_tokens() * packed_bytes {
return Err(PlaidError::InvalidIndex(format!(
"doc_residual_bytes length {} is not tokens ({}) × packed_bytes ({packed_bytes})",
index.doc_residual_bytes.len(),
index.num_tokens(),
)));
}
write_u32_slice(w, &index.doc_centroid_ids)?;
w.write_all(&index.doc_residual_bytes)?;
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}; run `docbert rebuild` to regenerate the \
index from the stored embeddings",
)));
}
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 total_tokens: usize = token_counts.iter().map(|&c| c as usize).sum();
let doc_centroid_ids = read_u32_vec(r, total_tokens)?;
let doc_residual_bytes = read_u8_vec(r, total_tokens * packed_bytes)?;
if let Some(&bad) = doc_centroid_ids
.iter()
.find(|&&id| (id as usize) >= k_centroids)
{
return Err(PlaidError::InvalidIndex(format!(
"centroid_id {bad} out of range 0..{k_centroids}",
)));
}
let mut doc_offsets: Vec<usize> = Vec::with_capacity(n_documents + 1);
doc_offsets.push(0);
let mut acc = 0usize;
for &count in &token_counts {
acc += count as usize;
doc_offsets.push(acc);
}
let ivf = build_inverted_file_from_flat(
&doc_centroid_ids,
&doc_offsets,
k_centroids,
);
let codec = ResidualCodec {
nbits,
dim,
centroids,
bucket_cutoffs,
bucket_weights,
};
codec.validate()?;
Ok(Index {
params,
codec,
doc_ids,
doc_centroid_ids,
doc_residual_bytes,
doc_offsets,
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_u8_vec<R: Read>(r: &mut R, n: usize) -> Result<Vec<u8>> {
let mut out = vec![0u8; n];
r.read_exact(&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.num_documents(), index.num_documents());
for i in 0..index.num_documents() {
assert_eq!(loaded.doc_tokens_vec(i), index.doc_tokens_vec(i));
}
}
#[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_token_count(loaded.num_documents() - 1), 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_versions_with_rebuild_instructions() {
let tmp = tempfile::tempdir().unwrap();
for version in [1u32, 2] {
let path = tmp.path().join(format!("legacy_v{version}.plaid"));
let mut buf = Vec::new();
buf.extend_from_slice(MAGIC);
buf.extend_from_slice(&version.to_le_bytes());
std::fs::write(&path, &buf).unwrap();
let err = load(&path).unwrap_err();
let PlaidError::InvalidIndex(msg) = &err else {
panic!("expected InvalidIndex for v{version}, got {err:?}");
};
assert!(msg.contains("version"), "v{version}: {msg}");
assert!(msg.contains("docbert rebuild"), "v{version}: {msg}");
}
}
}