1use std::{
32 fs::File,
33 io::{BufReader, BufWriter, Read, Write},
34 path::Path,
35};
36
37use crate::{
38 PlaidError,
39 Result,
40 codec::{ResidualCodec, packed_bytes_per_vector},
41 index::{Index, IndexParams, build_inverted_file_from_flat},
42};
43
44const MAGIC: &[u8; 8] = b"PLAIDIDX";
45const FORMAT_VERSION: u32 = 3;
58
59pub fn save(index: &Index, path: &Path) -> Result<()> {
72 let file = File::create(path)?;
73 let mut writer = BufWriter::new(file);
74 write_index(index, &mut writer)?;
75 writer.flush()?;
76 Ok(())
77}
78
79pub fn load(path: &Path) -> Result<Index> {
90 let file = File::open(path)?;
91 let mut reader = BufReader::new(file);
92 read_index(&mut reader)
93}
94
95fn write_index<W: Write>(index: &Index, w: &mut W) -> Result<()> {
96 w.write_all(MAGIC)?;
97 write_u32(w, FORMAT_VERSION)?;
98
99 let params = &index.params;
100 write_u32(w, params.dim as u32)?;
101 write_u32(w, params.nbits)?;
102 write_u32(w, params.k_centroids as u32)?;
103 write_u32(w, params.max_kmeans_iters as u32)?;
104 write_u64(w, index.doc_ids.len() as u64)?;
105
106 write_f32_slice(w, &index.codec.centroids)?;
107 write_f32_slice(w, &index.codec.bucket_cutoffs)?;
108 write_f32_slice(w, &index.codec.bucket_weights)?;
109
110 write_u64_slice(w, &index.doc_ids)?;
111
112 let token_counts: Vec<u32> = (0..index.num_documents())
113 .map(|i| index.doc_token_count(i) as u32)
114 .collect();
115 write_u32_slice(w, &token_counts)?;
116
117 let packed_bytes = packed_bytes_per_vector(params.dim, params.nbits);
118 if index.doc_residual_bytes.len() != index.num_tokens() * packed_bytes {
119 return Err(PlaidError::InvalidIndex(format!(
120 "doc_residual_bytes length {} is not tokens ({}) × packed_bytes ({packed_bytes})",
121 index.doc_residual_bytes.len(),
122 index.num_tokens(),
123 )));
124 }
125
126 write_u32_slice(w, &index.doc_centroid_ids)?;
127 w.write_all(&index.doc_residual_bytes)?;
128
129 Ok(())
130}
131
132fn read_index<R: Read>(r: &mut R) -> Result<Index> {
133 let mut magic = [0u8; 8];
134 r.read_exact(&mut magic)?;
135 if &magic != MAGIC {
136 return Err(PlaidError::InvalidIndex(
137 "not a docbert-plaid index (magic bytes mismatch)".into(),
138 ));
139 }
140
141 let version = read_u32(r)?;
142 if version != FORMAT_VERSION {
143 return Err(PlaidError::InvalidIndex(format!(
144 "unsupported plaid index version {version}, expected \
145 {FORMAT_VERSION}; run `docbert rebuild` to regenerate the \
146 index from the stored embeddings",
147 )));
148 }
149
150 let dim = read_u32(r)? as usize;
151 let nbits = read_u32(r)?;
152 let k_centroids = read_u32(r)? as usize;
153 let max_kmeans_iters = read_u32(r)? as usize;
154 let n_documents = read_u64(r)? as usize;
155
156 if dim == 0 || k_centroids == 0 || !matches!(nbits, 1 | 2 | 4 | 8) {
157 return Err(PlaidError::InvalidIndex(
158 "plaid index header has invalid dim/k_centroids/nbits".into(),
159 ));
160 }
161
162 let params = IndexParams {
163 dim,
164 nbits,
165 k_centroids,
166 max_kmeans_iters,
167 };
168
169 let centroids = read_f32_vec(r, k_centroids * dim)?;
170 let num_buckets = 1usize << nbits;
171 let bucket_cutoffs = read_f32_vec(r, num_buckets - 1)?;
172 let bucket_weights = read_f32_vec(r, num_buckets)?;
173
174 let doc_ids = read_u64_vec(r, n_documents)?;
175 let token_counts = read_u32_vec(r, n_documents)?;
176 let packed_bytes = packed_bytes_per_vector(dim, nbits);
177
178 let total_tokens: usize = token_counts.iter().map(|&c| c as usize).sum();
179
180 let doc_centroid_ids = read_u32_vec(r, total_tokens)?;
182 let doc_residual_bytes = read_u8_vec(r, total_tokens * packed_bytes)?;
183
184 if let Some(&bad) = doc_centroid_ids
187 .iter()
188 .find(|&&id| (id as usize) >= k_centroids)
189 {
190 return Err(PlaidError::InvalidIndex(format!(
191 "centroid_id {bad} out of range 0..{k_centroids}",
192 )));
193 }
194
195 let mut doc_offsets: Vec<usize> = Vec::with_capacity(n_documents + 1);
196 doc_offsets.push(0);
197 let mut acc = 0usize;
198 for &count in &token_counts {
199 acc += count as usize;
200 doc_offsets.push(acc);
201 }
202
203 let ivf = build_inverted_file_from_flat(
204 &doc_centroid_ids,
205 &doc_offsets,
206 k_centroids,
207 );
208
209 let codec = ResidualCodec {
210 nbits,
211 dim,
212 centroids,
213 bucket_cutoffs,
214 bucket_weights,
215 };
216 codec.validate()?;
217
218 Ok(Index {
219 params,
220 codec,
221 doc_ids,
222 doc_centroid_ids,
223 doc_residual_bytes,
224 doc_offsets,
225 ivf,
226 })
227}
228
229fn write_u32<W: Write>(w: &mut W, v: u32) -> Result<()> {
230 w.write_all(&v.to_le_bytes())?;
231 Ok(())
232}
233
234fn write_u64<W: Write>(w: &mut W, v: u64) -> Result<()> {
235 w.write_all(&v.to_le_bytes())?;
236 Ok(())
237}
238
239fn write_f32_slice<W: Write>(w: &mut W, slice: &[f32]) -> Result<()> {
240 w.write_all(bytemuck::cast_slice(slice))?;
241 Ok(())
242}
243
244fn write_u32_slice<W: Write>(w: &mut W, slice: &[u32]) -> Result<()> {
245 w.write_all(bytemuck::cast_slice(slice))?;
246 Ok(())
247}
248
249fn write_u64_slice<W: Write>(w: &mut W, slice: &[u64]) -> Result<()> {
250 w.write_all(bytemuck::cast_slice(slice))?;
251 Ok(())
252}
253
254fn read_u32<R: Read>(r: &mut R) -> Result<u32> {
255 let mut buf = [0u8; 4];
256 r.read_exact(&mut buf)?;
257 Ok(u32::from_le_bytes(buf))
258}
259
260fn read_u64<R: Read>(r: &mut R) -> Result<u64> {
261 let mut buf = [0u8; 8];
262 r.read_exact(&mut buf)?;
263 Ok(u64::from_le_bytes(buf))
264}
265
266fn read_f32_vec<R: Read>(r: &mut R, n: usize) -> Result<Vec<f32>> {
267 let mut out = vec![0.0f32; n];
268 r.read_exact(bytemuck::cast_slice_mut(&mut out))?;
269 Ok(out)
270}
271
272fn read_u8_vec<R: Read>(r: &mut R, n: usize) -> Result<Vec<u8>> {
273 let mut out = vec![0u8; n];
274 r.read_exact(&mut out)?;
275 Ok(out)
276}
277
278fn read_u32_vec<R: Read>(r: &mut R, n: usize) -> Result<Vec<u32>> {
279 let mut out = vec![0u32; n];
280 r.read_exact(bytemuck::cast_slice_mut(&mut out))?;
281 Ok(out)
282}
283
284fn read_u64_vec<R: Read>(r: &mut R, n: usize) -> Result<Vec<u64>> {
285 let mut out = vec![0u64; n];
286 r.read_exact(bytemuck::cast_slice_mut(&mut out))?;
287 Ok(out)
288}
289
290#[cfg(test)]
291mod tests {
292 use super::*;
293 use crate::index::{DocumentTokens, build_index};
294
295 fn small_corpus() -> Vec<DocumentTokens> {
296 vec![
297 DocumentTokens {
298 doc_id: 10,
299 tokens: vec![0.0, 0.0, 0.1, 0.2, -0.1, 0.1],
300 n_tokens: 3,
301 },
302 DocumentTokens {
303 doc_id: 20,
304 tokens: vec![10.0, 10.0, 10.2, 9.9, 9.8, 10.1],
305 n_tokens: 3,
306 },
307 DocumentTokens {
308 doc_id: 30,
309 tokens: vec![0.3, -0.2, 9.7, 10.2],
310 n_tokens: 2,
311 },
312 ]
313 }
314
315 fn default_params() -> IndexParams {
316 IndexParams {
317 dim: 2,
318 nbits: 2,
319 k_centroids: 2,
320 max_kmeans_iters: 50,
321 }
322 }
323
324 #[test]
325 fn round_trip_preserves_codec_parameters_and_doc_ids() {
326 let tmp = tempfile::tempdir().unwrap();
327 let path = tmp.path().join("index.plaid");
328 let index = build_index(&small_corpus(), default_params()).unwrap();
329
330 save(&index, &path).unwrap();
331 let loaded = load(&path).unwrap();
332
333 assert_eq!(loaded.params.dim, index.params.dim);
334 assert_eq!(loaded.params.nbits, index.params.nbits);
335 assert_eq!(loaded.params.k_centroids, index.params.k_centroids);
336 assert_eq!(
337 loaded.params.max_kmeans_iters,
338 index.params.max_kmeans_iters
339 );
340 assert_eq!(loaded.doc_ids, index.doc_ids);
341 }
342
343 #[test]
344 fn round_trip_preserves_codec_tables_byte_for_byte() {
345 let tmp = tempfile::tempdir().unwrap();
346 let path = tmp.path().join("index.plaid");
347 let index = build_index(&small_corpus(), default_params()).unwrap();
348
349 save(&index, &path).unwrap();
350 let loaded = load(&path).unwrap();
351
352 assert_eq!(loaded.codec.centroids, index.codec.centroids);
353 assert_eq!(loaded.codec.bucket_cutoffs, index.codec.bucket_cutoffs);
354 assert_eq!(loaded.codec.bucket_weights, index.codec.bucket_weights);
355 }
356
357 #[test]
358 fn round_trip_preserves_encoded_tokens() {
359 let tmp = tempfile::tempdir().unwrap();
360 let path = tmp.path().join("index.plaid");
361 let index = build_index(&small_corpus(), default_params()).unwrap();
362
363 save(&index, &path).unwrap();
364 let loaded = load(&path).unwrap();
365
366 assert_eq!(loaded.num_documents(), index.num_documents());
367 for i in 0..index.num_documents() {
368 assert_eq!(loaded.doc_tokens_vec(i), index.doc_tokens_vec(i));
369 }
370 }
371
372 #[test]
373 fn round_trip_rebuilds_the_inverted_file() {
374 let tmp = tempfile::tempdir().unwrap();
375 let path = tmp.path().join("index.plaid");
376 let index = build_index(&small_corpus(), default_params()).unwrap();
377
378 save(&index, &path).unwrap();
379 let loaded = load(&path).unwrap();
380
381 assert_eq!(loaded.ivf.num_centroids(), index.ivf.num_centroids());
382 assert_eq!(
383 loaded.ivf.total_doc_postings(),
384 index.ivf.total_doc_postings(),
385 );
386 for c in 0..index.ivf.num_centroids() {
387 let want = index.ivf.docs_for_centroid(c);
388 let got = loaded.ivf.docs_for_centroid(c);
389 assert_eq!(got, want, "IVF list for centroid {c} differs");
390 }
391 }
392
393 #[test]
394 fn round_trip_preserves_search_results_exactly() {
395 let tmp = tempfile::tempdir().unwrap();
396 let path = tmp.path().join("index.plaid");
397 let index = build_index(&small_corpus(), default_params()).unwrap();
398 save(&index, &path).unwrap();
399 let loaded = load(&path).unwrap();
400
401 let query = [0.05f32, 0.1, 9.9, 10.1];
402 let params = crate::search::SearchParams {
403 top_k: 3,
404 n_probe: 2,
405 n_candidate_docs: None,
406 centroid_score_threshold: None,
407 };
408
409 let a = crate::search::search(&index, &query, params).unwrap();
410 let b = crate::search::search(&loaded, &query, params).unwrap();
411 assert_eq!(a, b);
412 }
413
414 #[test]
415 fn round_trip_handles_empty_documents() {
416 let tmp = tempfile::tempdir().unwrap();
417 let path = tmp.path().join("index.plaid");
418 let mut docs = small_corpus();
419 docs.push(DocumentTokens {
420 doc_id: 999,
421 tokens: vec![],
422 n_tokens: 0,
423 });
424 let index = build_index(&docs, default_params()).unwrap();
425
426 save(&index, &path).unwrap();
427 let loaded = load(&path).unwrap();
428
429 assert_eq!(loaded.doc_ids, index.doc_ids);
430 assert_eq!(loaded.doc_token_count(loaded.num_documents() - 1), 0);
431 }
432
433 #[test]
434 fn load_rejects_files_with_wrong_magic() {
435 let tmp = tempfile::tempdir().unwrap();
436 let path = tmp.path().join("bogus.plaid");
437 std::fs::write(&path, b"NOTPLAID").unwrap();
438
439 let err = load(&path).unwrap_err();
440 assert!(
441 matches!(err, PlaidError::InvalidIndex(ref m) if m.contains("magic")),
442 "expected InvalidIndex with magic message, got {err:?}",
443 );
444 }
445
446 #[test]
447 fn load_rejects_unknown_format_version() {
448 let tmp = tempfile::tempdir().unwrap();
449 let path = tmp.path().join("future.plaid");
450 let mut buf = Vec::new();
451 buf.extend_from_slice(MAGIC);
452 buf.extend_from_slice(&999u32.to_le_bytes());
453 std::fs::write(&path, &buf).unwrap();
454
455 let err = load(&path).unwrap_err();
456 assert!(
457 matches!(err, PlaidError::InvalidIndex(ref m) if m.contains("version")),
458 "expected InvalidIndex with version message, got {err:?}",
459 );
460 }
461
462 #[test]
463 fn load_rejects_legacy_versions_with_rebuild_instructions() {
464 let tmp = tempfile::tempdir().unwrap();
467 for version in [1u32, 2] {
468 let path = tmp.path().join(format!("legacy_v{version}.plaid"));
469 let mut buf = Vec::new();
470 buf.extend_from_slice(MAGIC);
471 buf.extend_from_slice(&version.to_le_bytes());
472 std::fs::write(&path, &buf).unwrap();
473
474 let err = load(&path).unwrap_err();
475 let PlaidError::InvalidIndex(msg) = &err else {
476 panic!("expected InvalidIndex for v{version}, got {err:?}");
477 };
478 assert!(msg.contains("version"), "v{version}: {msg}");
479 assert!(msg.contains("docbert rebuild"), "v{version}: {msg}");
480 }
481 }
482}