Skip to main content

docbert_plaid/
persistence.rs

1//! Save and load a built [`Index`] to/from disk.
2//!
3//! The on-disk format is a single little-endian binary file with a small
4//! header followed by a sequence of plain f32/u32/u64 blobs. The layout
5//! is deliberately boring — there's no compression, no extra framing,
6//! and no schema evolution today; we can layer any of that on top later
7//! without changing the semantic API.
8//!
9//! Layout, every field little-endian:
10//!
11//! ```text
12//! magic           : 8 bytes, b"PLAIDIDX"
13//! version         : u32     (currently 1)
14//! dim             : u32
15//! nbits           : u32
16//! k_centroids    : u32
17//! max_kmeans_iters: u32
18//! n_documents     : u64
19//! centroids       : dim * k_centroids     f32
20//! bucket_cutoffs  : 2^nbits - 1            f32
21//! bucket_weights  : 2^nbits                f32
22//! doc_ids         : n_documents            u64
23//! token_counts    : n_documents            u32  (tokens per document)
24//! encoded_tokens  : for each token, u32 centroid_id then `dim` u8 codes
25//! ```
26//!
27//! The inverted file is *not* persisted — it's derivable from the
28//! encoded tokens in O(n_tokens) time on load and storing it would just
29//! duplicate state that's already on disk elsewhere.
30
31use std::{
32    fs::File,
33    io::{BufReader, BufWriter, Read, Write},
34    path::Path,
35};
36
37use crate::{
38    PlaidError,
39    Result,
40    codec::{EncodedVector, ResidualCodec, packed_bytes_per_vector},
41    index::{Index, IndexParams, build_inverted_file},
42};
43
44const MAGIC: &[u8; 8] = b"PLAIDIDX";
45/// Format versions in the wild:
46///
47/// - `1`: unpacked residual codes, one byte per residual dimension.
48/// - `2`: LSB-first bit-packed codes at `nbits ∈ {1, 2, 4, 8}`;
49///   `(dim * nbits) / 8` bytes per token.
50const FORMAT_VERSION: u32 = 2;
51
52/// Write `index` to `path`, creating or truncating the file as needed.
53///
54/// # Errors
55///
56/// Returns [`PlaidError::Io`] for any underlying I/O failure, or
57/// [`PlaidError::InvalidIndex`] if an encoded token's `codes` buffer
58/// doesn't match the codec's advertised packed length (a caller bug
59/// worth flagging loudly on write rather than producing a file that
60/// would fail to load back).
61///
62/// [`PlaidError::Io`]: crate::PlaidError::Io
63/// [`PlaidError::InvalidIndex`]: crate::PlaidError::InvalidIndex
64pub fn save(index: &Index, path: &Path) -> Result<()> {
65    let file = File::create(path)?;
66    let mut writer = BufWriter::new(file);
67    write_index(index, &mut writer)?;
68    writer.flush()?;
69    Ok(())
70}
71
72/// Read an [`Index`] back from `path`.
73///
74/// # Errors
75///
76/// Returns [`PlaidError::Io`] for any read failure, or
77/// [`PlaidError::InvalidIndex`] if the file's magic, version, or
78/// header fields don't match this crate's format.
79///
80/// [`PlaidError::Io`]: crate::PlaidError::Io
81/// [`PlaidError::InvalidIndex`]: crate::PlaidError::InvalidIndex
82pub fn load(path: &Path) -> Result<Index> {
83    let file = File::open(path)?;
84    let mut reader = BufReader::new(file);
85    read_index(&mut reader)
86}
87
88fn write_index<W: Write>(index: &Index, w: &mut W) -> Result<()> {
89    w.write_all(MAGIC)?;
90    write_u32(w, FORMAT_VERSION)?;
91
92    let params = &index.params;
93    write_u32(w, params.dim as u32)?;
94    write_u32(w, params.nbits)?;
95    write_u32(w, params.k_centroids as u32)?;
96    write_u32(w, params.max_kmeans_iters as u32)?;
97    write_u64(w, index.doc_ids.len() as u64)?;
98
99    write_f32_slice(w, &index.codec.centroids)?;
100    write_f32_slice(w, &index.codec.bucket_cutoffs)?;
101    write_f32_slice(w, &index.codec.bucket_weights)?;
102
103    write_u64_slice(w, &index.doc_ids)?;
104
105    let token_counts: Vec<u32> = index
106        .doc_tokens
107        .iter()
108        .map(|tokens| tokens.len() as u32)
109        .collect();
110    write_u32_slice(w, &token_counts)?;
111
112    let packed_bytes = packed_bytes_per_vector(params.dim, params.nbits);
113    for encoded_doc in &index.doc_tokens {
114        for ev in encoded_doc {
115            write_u32(w, ev.centroid_id)?;
116            if ev.codes.len() != packed_bytes {
117                return Err(PlaidError::InvalidIndex(format!(
118                    "encoded token has {} packed bytes but codec expects {packed_bytes}",
119                    ev.codes.len(),
120                )));
121            }
122            w.write_all(&ev.codes)?;
123        }
124    }
125
126    Ok(())
127}
128
129fn read_index<R: Read>(r: &mut R) -> Result<Index> {
130    let mut magic = [0u8; 8];
131    r.read_exact(&mut magic)?;
132    if &magic != MAGIC {
133        return Err(PlaidError::InvalidIndex(
134            "not a docbert-plaid index (magic bytes mismatch)".into(),
135        ));
136    }
137
138    let version = read_u32(r)?;
139    if version != FORMAT_VERSION {
140        return Err(PlaidError::InvalidIndex(format!(
141            "unsupported plaid index version {version}, expected {FORMAT_VERSION}",
142        )));
143    }
144
145    let dim = read_u32(r)? as usize;
146    let nbits = read_u32(r)?;
147    let k_centroids = read_u32(r)? as usize;
148    let max_kmeans_iters = read_u32(r)? as usize;
149    let n_documents = read_u64(r)? as usize;
150
151    if dim == 0 || k_centroids == 0 || !matches!(nbits, 1 | 2 | 4 | 8) {
152        return Err(PlaidError::InvalidIndex(
153            "plaid index header has invalid dim/k_centroids/nbits".into(),
154        ));
155    }
156
157    let params = IndexParams {
158        dim,
159        nbits,
160        k_centroids,
161        max_kmeans_iters,
162    };
163
164    let centroids = read_f32_vec(r, k_centroids * dim)?;
165    let num_buckets = 1usize << nbits;
166    let bucket_cutoffs = read_f32_vec(r, num_buckets - 1)?;
167    let bucket_weights = read_f32_vec(r, num_buckets)?;
168
169    let doc_ids = read_u64_vec(r, n_documents)?;
170    let token_counts = read_u32_vec(r, n_documents)?;
171    let packed_bytes = packed_bytes_per_vector(dim, nbits);
172
173    let mut doc_tokens: Vec<Vec<EncodedVector>> =
174        Vec::with_capacity(n_documents);
175    for count in token_counts.iter() {
176        let mut encoded_doc = Vec::with_capacity(*count as usize);
177        for _ in 0..*count {
178            let centroid_id = read_u32(r)?;
179            if (centroid_id as usize) >= k_centroids {
180                return Err(PlaidError::InvalidIndex(format!(
181                    "centroid_id {centroid_id} out of range 0..{k_centroids}",
182                )));
183            }
184            let mut codes = vec![0u8; packed_bytes];
185            r.read_exact(&mut codes)?;
186            encoded_doc.push(EncodedVector { centroid_id, codes });
187        }
188        doc_tokens.push(encoded_doc);
189    }
190    let ivf = build_inverted_file(&doc_tokens, k_centroids);
191
192    let codec = ResidualCodec {
193        nbits,
194        dim,
195        centroids,
196        bucket_cutoffs,
197        bucket_weights,
198    };
199    codec.validate()?;
200
201    Ok(Index {
202        params,
203        codec,
204        doc_ids,
205        doc_tokens,
206        ivf,
207    })
208}
209
210fn write_u32<W: Write>(w: &mut W, v: u32) -> Result<()> {
211    w.write_all(&v.to_le_bytes())?;
212    Ok(())
213}
214
215fn write_u64<W: Write>(w: &mut W, v: u64) -> Result<()> {
216    w.write_all(&v.to_le_bytes())?;
217    Ok(())
218}
219
220fn write_f32_slice<W: Write>(w: &mut W, slice: &[f32]) -> Result<()> {
221    w.write_all(bytemuck::cast_slice(slice))?;
222    Ok(())
223}
224
225fn write_u32_slice<W: Write>(w: &mut W, slice: &[u32]) -> Result<()> {
226    w.write_all(bytemuck::cast_slice(slice))?;
227    Ok(())
228}
229
230fn write_u64_slice<W: Write>(w: &mut W, slice: &[u64]) -> Result<()> {
231    w.write_all(bytemuck::cast_slice(slice))?;
232    Ok(())
233}
234
235fn read_u32<R: Read>(r: &mut R) -> Result<u32> {
236    let mut buf = [0u8; 4];
237    r.read_exact(&mut buf)?;
238    Ok(u32::from_le_bytes(buf))
239}
240
241fn read_u64<R: Read>(r: &mut R) -> Result<u64> {
242    let mut buf = [0u8; 8];
243    r.read_exact(&mut buf)?;
244    Ok(u64::from_le_bytes(buf))
245}
246
247fn read_f32_vec<R: Read>(r: &mut R, n: usize) -> Result<Vec<f32>> {
248    let mut out = vec![0.0f32; n];
249    r.read_exact(bytemuck::cast_slice_mut(&mut out))?;
250    Ok(out)
251}
252
253fn read_u32_vec<R: Read>(r: &mut R, n: usize) -> Result<Vec<u32>> {
254    let mut out = vec![0u32; n];
255    r.read_exact(bytemuck::cast_slice_mut(&mut out))?;
256    Ok(out)
257}
258
259fn read_u64_vec<R: Read>(r: &mut R, n: usize) -> Result<Vec<u64>> {
260    let mut out = vec![0u64; n];
261    r.read_exact(bytemuck::cast_slice_mut(&mut out))?;
262    Ok(out)
263}
264
265#[cfg(test)]
266mod tests {
267    use super::*;
268    use crate::index::{DocumentTokens, build_index};
269
270    fn small_corpus() -> Vec<DocumentTokens> {
271        vec![
272            DocumentTokens {
273                doc_id: 10,
274                tokens: vec![0.0, 0.0, 0.1, 0.2, -0.1, 0.1],
275                n_tokens: 3,
276            },
277            DocumentTokens {
278                doc_id: 20,
279                tokens: vec![10.0, 10.0, 10.2, 9.9, 9.8, 10.1],
280                n_tokens: 3,
281            },
282            DocumentTokens {
283                doc_id: 30,
284                tokens: vec![0.3, -0.2, 9.7, 10.2],
285                n_tokens: 2,
286            },
287        ]
288    }
289
290    fn default_params() -> IndexParams {
291        IndexParams {
292            dim: 2,
293            nbits: 2,
294            k_centroids: 2,
295            max_kmeans_iters: 50,
296        }
297    }
298
299    #[test]
300    fn round_trip_preserves_codec_parameters_and_doc_ids() {
301        let tmp = tempfile::tempdir().unwrap();
302        let path = tmp.path().join("index.plaid");
303        let index = build_index(&small_corpus(), default_params()).unwrap();
304
305        save(&index, &path).unwrap();
306        let loaded = load(&path).unwrap();
307
308        assert_eq!(loaded.params.dim, index.params.dim);
309        assert_eq!(loaded.params.nbits, index.params.nbits);
310        assert_eq!(loaded.params.k_centroids, index.params.k_centroids);
311        assert_eq!(
312            loaded.params.max_kmeans_iters,
313            index.params.max_kmeans_iters
314        );
315        assert_eq!(loaded.doc_ids, index.doc_ids);
316    }
317
318    #[test]
319    fn round_trip_preserves_codec_tables_byte_for_byte() {
320        let tmp = tempfile::tempdir().unwrap();
321        let path = tmp.path().join("index.plaid");
322        let index = build_index(&small_corpus(), default_params()).unwrap();
323
324        save(&index, &path).unwrap();
325        let loaded = load(&path).unwrap();
326
327        assert_eq!(loaded.codec.centroids, index.codec.centroids);
328        assert_eq!(loaded.codec.bucket_cutoffs, index.codec.bucket_cutoffs);
329        assert_eq!(loaded.codec.bucket_weights, index.codec.bucket_weights);
330    }
331
332    #[test]
333    fn round_trip_preserves_encoded_tokens() {
334        let tmp = tempfile::tempdir().unwrap();
335        let path = tmp.path().join("index.plaid");
336        let index = build_index(&small_corpus(), default_params()).unwrap();
337
338        save(&index, &path).unwrap();
339        let loaded = load(&path).unwrap();
340
341        assert_eq!(loaded.doc_tokens.len(), index.doc_tokens.len());
342        for (a, b) in loaded.doc_tokens.iter().zip(index.doc_tokens.iter()) {
343            assert_eq!(a, b);
344        }
345    }
346
347    #[test]
348    fn round_trip_rebuilds_the_inverted_file() {
349        let tmp = tempfile::tempdir().unwrap();
350        let path = tmp.path().join("index.plaid");
351        let index = build_index(&small_corpus(), default_params()).unwrap();
352
353        save(&index, &path).unwrap();
354        let loaded = load(&path).unwrap();
355
356        assert_eq!(loaded.ivf.num_centroids(), index.ivf.num_centroids());
357        assert_eq!(
358            loaded.ivf.total_doc_postings(),
359            index.ivf.total_doc_postings(),
360        );
361        for c in 0..index.ivf.num_centroids() {
362            let want = index.ivf.docs_for_centroid(c);
363            let got = loaded.ivf.docs_for_centroid(c);
364            assert_eq!(got, want, "IVF list for centroid {c} differs");
365        }
366    }
367
368    #[test]
369    fn round_trip_preserves_search_results_exactly() {
370        let tmp = tempfile::tempdir().unwrap();
371        let path = tmp.path().join("index.plaid");
372        let index = build_index(&small_corpus(), default_params()).unwrap();
373        save(&index, &path).unwrap();
374        let loaded = load(&path).unwrap();
375
376        let query = [0.05f32, 0.1, 9.9, 10.1];
377        let params = crate::search::SearchParams {
378            top_k: 3,
379            n_probe: 2,
380            n_candidate_docs: None,
381            centroid_score_threshold: None,
382        };
383
384        let a = crate::search::search(&index, &query, params).unwrap();
385        let b = crate::search::search(&loaded, &query, params).unwrap();
386        assert_eq!(a, b);
387    }
388
389    #[test]
390    fn round_trip_handles_empty_documents() {
391        let tmp = tempfile::tempdir().unwrap();
392        let path = tmp.path().join("index.plaid");
393        let mut docs = small_corpus();
394        docs.push(DocumentTokens {
395            doc_id: 999,
396            tokens: vec![],
397            n_tokens: 0,
398        });
399        let index = build_index(&docs, default_params()).unwrap();
400
401        save(&index, &path).unwrap();
402        let loaded = load(&path).unwrap();
403
404        assert_eq!(loaded.doc_ids, index.doc_ids);
405        assert_eq!(loaded.doc_tokens.last().unwrap().len(), 0);
406    }
407
408    #[test]
409    fn load_rejects_files_with_wrong_magic() {
410        let tmp = tempfile::tempdir().unwrap();
411        let path = tmp.path().join("bogus.plaid");
412        std::fs::write(&path, b"NOTPLAID").unwrap();
413
414        let err = load(&path).unwrap_err();
415        assert!(
416            matches!(err, PlaidError::InvalidIndex(ref m) if m.contains("magic")),
417            "expected InvalidIndex with magic message, got {err:?}",
418        );
419    }
420
421    #[test]
422    fn load_rejects_unknown_format_version() {
423        let tmp = tempfile::tempdir().unwrap();
424        let path = tmp.path().join("future.plaid");
425        let mut buf = Vec::new();
426        buf.extend_from_slice(MAGIC);
427        buf.extend_from_slice(&999u32.to_le_bytes());
428        std::fs::write(&path, &buf).unwrap();
429
430        let err = load(&path).unwrap_err();
431        assert!(
432            matches!(err, PlaidError::InvalidIndex(ref m) if m.contains("version")),
433            "expected InvalidIndex with version message, got {err:?}",
434        );
435    }
436
437    #[test]
438    fn load_rejects_legacy_unpacked_format_version_one() {
439        // Version 1 (pre-packing) indexes must be rebuilt. The loader
440        // should reject them with a clear version mismatch.
441        let tmp = tempfile::tempdir().unwrap();
442        let path = tmp.path().join("legacy.plaid");
443        let mut buf = Vec::new();
444        buf.extend_from_slice(MAGIC);
445        buf.extend_from_slice(&1u32.to_le_bytes());
446        std::fs::write(&path, &buf).unwrap();
447
448        let err = load(&path).unwrap_err();
449        assert!(
450            matches!(err, PlaidError::InvalidIndex(ref m) if m.contains("version")),
451            "expected InvalidIndex with version message, got {err:?}",
452        );
453    }
454}