1use 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";
45const FORMAT_VERSION: u32 = 2;
51
52pub 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
72pub 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 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}