Skip to main content

hermes_core/index/
vector_builder.rs

1//! Vector index building for IndexWriter
2//!
3//! Training is **manual-only** — decoupled from commit.
4//! Call `build_vector_index()` explicitly when ready.
5//! ANN indexes are built naturally during subsequent merges.
6
7use std::io::Write;
8use std::sync::Arc;
9
10use rustc_hash::FxHashMap;
11
12use crate::directories::DirectoryWriter;
13use crate::dsl::{DenseVectorConfig, Field, FieldType, VectorIndexType};
14use crate::error::{Error, Result};
15use crate::segment::{SegmentId, SegmentReader};
16
17use super::IndexWriter;
18
19/// Maximum supported IVF centroid count. Query-side `nprobe` and serialized
20/// cluster identifiers use the same practical bound.
21const MAX_IVF_CLUSTERS: usize = 4096;
22
23struct TrainedFieldUpdate {
24    field_id: u32,
25    index_type: VectorIndexType,
26    vector_count: usize,
27    num_clusters: usize,
28    centroids_file: String,
29    codebook_file: Option<String>,
30}
31
32/// Write adapter that rejects an artifact before its serialized form exceeds
33/// the same bound enforced by the loader. Encoding directly through this
34/// adapter avoids materializing a second, potentially hundreds-of-megabytes
35/// copy of the trained structure.
36struct SizeLimitedWriter<'a, W: Write + ?Sized> {
37    inner: &'a mut W,
38    written: usize,
39    limit: usize,
40}
41
42impl<'a, W: Write + ?Sized> SizeLimitedWriter<'a, W> {
43    fn new(inner: &'a mut W, limit: usize) -> Self {
44        Self {
45            inner,
46            written: 0,
47            limit,
48        }
49    }
50}
51
52impl<W: Write + ?Sized> Write for SizeLimitedWriter<'_, W> {
53    fn write(&mut self, buffer: &[u8]) -> std::io::Result<usize> {
54        let next_size = self
55            .written
56            .checked_add(buffer.len())
57            .ok_or_else(|| std::io::Error::other("trained artifact size overflow"))?;
58        if next_size > self.limit {
59            return Err(std::io::Error::new(
60                std::io::ErrorKind::InvalidData,
61                format!(
62                    "trained artifact exceeds the {}-byte safety limit",
63                    self.limit
64                ),
65            ));
66        }
67        let written = self.inner.write(buffer)?;
68        self.written += written;
69        Ok(written)
70    }
71
72    fn flush(&mut self) -> std::io::Result<()> {
73        self.inner.flush()
74    }
75}
76
77fn validate_explicit_ivf_num_clusters(config: &DenseVectorConfig) -> Result<()> {
78    match config.num_clusters {
79        Some(0) => Err(Error::Schema(
80            "dense vector num_clusters must be at least 1".to_string(),
81        )),
82        Some(value) if value > MAX_IVF_CLUSTERS => Err(Error::Schema(format!(
83            "dense vector num_clusters must not exceed {MAX_IVF_CLUSTERS}, got {value}"
84        ))),
85        _ => Ok(()),
86    }
87}
88
89/// Validate the configured centroid count and cap it to the training sample.
90///
91/// Corpus size drives the automatic heuristic, but training cannot produce
92/// more distinct centroids than the number of sampled vectors. Keeping this
93/// decision here avoids relying on a panic-prone, implicit clamp inside the
94/// trainer and gives callers a schema error for invalid explicit values.
95fn effective_ivf_num_clusters(
96    config: &DenseVectorConfig,
97    corpus_count: usize,
98    sample_count: usize,
99) -> Result<usize> {
100    if sample_count == 0 {
101        return Err(Error::Schema(
102            "cannot train an IVF vector index without sample vectors".to_string(),
103        ));
104    }
105
106    validate_explicit_ivf_num_clusters(config)?;
107    let requested = match config.num_clusters {
108        Some(value) => value,
109        None => config.optimal_num_clusters(corpus_count),
110    };
111
112    Ok(requested.min(sample_count))
113}
114
115impl<D: DirectoryWriter + 'static> IndexWriter<D> {
116    /// Train vector index from accumulated Flat vectors (manual, not auto-triggered).
117    ///
118    /// 1. Acquires a snapshot (segments safe to read)
119    /// 2. Collects vectors for training
120    /// 3. Trains centroids/codebooks
121    /// 4. Updates metadata (marks fields as Built)
122    /// 5. Publishes to ArcSwap — merges will use these automatically
123    ///
124    /// Existing flat segments get ANN during normal merges. No rebuild needed.
125    pub async fn build_vector_index(&self) -> Result<()> {
126        let dense_fields = self.get_dense_vector_fields();
127        if dense_fields.is_empty() {
128            log::info!("No dense vector fields configured for ANN indexing");
129            return Ok(());
130        }
131
132        let artifact_update = self.segment_manager.begin_vector_artifact_update().await?;
133        self.build_vector_index_locked(&dense_fields, &artifact_update)
134            .await
135    }
136
137    /// Build while the SegmentManager's artifact-update gate is held.
138    async fn build_vector_index_locked(
139        &self,
140        dense_fields: &[(Field, DenseVectorConfig)],
141        artifact_update: &crate::merge::VectorArtifactUpdateGuard,
142    ) -> Result<()> {
143        // Check which fields need building (skip already built)
144        let fields_to_build = self.get_fields_to_build(dense_fields).await;
145        if fields_to_build.is_empty() {
146            log::info!("All vector fields already built, skipping training");
147            return Ok(());
148        }
149
150        // Reject malformed explicit settings before opening segments or
151        // allocating the bounded training samples.
152        for (_, config) in &fields_to_build {
153            if config.uses_ivf() {
154                validate_explicit_ivf_num_clusters(config)?;
155            }
156        }
157
158        // Acquire snapshot — segments won't be deleted while we read them
159        let snapshot = self.segment_manager.acquire_snapshot().await;
160        let segment_ids = snapshot.segment_ids();
161        if segment_ids.is_empty() {
162            return Ok(());
163        }
164
165        // Collect vectors for training
166        let (all_vectors, total_vectors) = self
167            .collect_vectors_for_training(segment_ids, &fields_to_build)
168            .await?;
169
170        self.train_and_publish_fields(
171            &fields_to_build,
172            &all_vectors,
173            &total_vectors,
174            artifact_update,
175        )
176        .await
177    }
178
179    /// Train every requested field from pre-collected samples, then durably
180    /// publish the artifacts. If any field fails, all durable field states
181    /// remain Flat and the successfully written files are merely unreferenced
182    /// retry targets.
183    async fn train_and_publish_fields(
184        &self,
185        fields_to_build: &[(Field, DenseVectorConfig)],
186        all_vectors: &FxHashMap<u32, Vec<Vec<f32>>>,
187        total_vectors: &FxHashMap<u32, usize>,
188        artifact_update: &crate::merge::VectorArtifactUpdateGuard,
189    ) -> Result<()> {
190        let mut updates = Vec::with_capacity(fields_to_build.len());
191        for (field, config) in fields_to_build {
192            if let Some(update) = self
193                .train_field_index(*field, config, all_vectors, total_vectors)
194                .await?
195            {
196                updates.push(update);
197            }
198        }
199
200        if updates.is_empty() {
201            // Fail loud: training was explicitly requested and produced
202            // nothing — reporting success would leave callers believing the
203            // fields are Built.
204            let field_ids: Vec<u32> = fields_to_build.iter().map(|(field, _)| field.0).collect();
205            return Err(Error::Schema(format!(
206                "cannot train vector index: no training vectors were collected for \
207                 field(s) {field_ids:?}; commit documents containing these fields \
208                 before building"
209            )));
210        }
211
212        // Durable metadata and the complete validated ArcSwap set advance in a
213        // single cancellation-safe SegmentManager transaction.
214        self.segment_manager
215            .update_vector_metadata_and_publish(artifact_update, |meta| {
216                for update in &updates {
217                    meta.init_field(update.field_id, update.index_type);
218                    meta.mark_field_built(
219                        update.field_id,
220                        update.vector_count,
221                        update.num_clusters,
222                        update.centroids_file.clone(),
223                        update.codebook_file.clone(),
224                    );
225                }
226            })
227            .await?;
228
229        log::info!("Vector index training complete, ANN will be built during merges");
230
231        Ok(())
232    }
233
234    /// Rebuild vector index by retraining centroids/codebooks.
235    ///
236    /// Rebuilding a global artifact generation is only safe while every
237    /// committed segment is still flat. IVF/ScaNN segments embed the artifact
238    /// versions they were built with and cannot be interpreted by freshly
239    /// trained centroids/codebooks.
240    pub async fn rebuild_vector_index(&self) -> Result<()> {
241        let dense_fields = self.get_dense_vector_fields();
242        if dense_fields.is_empty() {
243            return Ok(());
244        }
245
246        // Raise the producer gate and drain operations that may already have
247        // captured the previous trained generation. New producers continue in
248        // flat mode until this guard drops.
249        let artifact_update = self.segment_manager.begin_vector_artifact_update().await?;
250        let snapshot = self.segment_manager.acquire_snapshot().await;
251        let field_ids: Vec<u32> = dense_fields.iter().map(|(field, _)| field.0).collect();
252        self.reject_rebuild_with_ann_segments(snapshot.segment_ids(), &field_ids)
253            .await?;
254
255        // Reject malformed explicit settings before collecting samples.
256        for (_, config) in &dense_fields {
257            if config.uses_ivf() {
258                validate_explicit_ivf_num_clusters(config)?;
259            }
260        }
261
262        // Collect the retraining samples BEFORE the durable Built -> Flat
263        // reset: a read failure (propagated by collect_vectors_for_training)
264        // or an empty sample for a Built field must not destructively
265        // downgrade the published artifact generation.
266        let (all_vectors, total_vectors) = self
267            .collect_vectors_for_training(snapshot.segment_ids(), &dense_fields)
268            .await?;
269        let built_fields: Vec<u32> = self
270            .segment_manager
271            .read_metadata(|meta| {
272                field_ids
273                    .iter()
274                    .filter(|field_id| meta.is_field_built(**field_id))
275                    .copied()
276                    .collect()
277            })
278            .await;
279        let starved_built: Vec<u32> = built_fields
280            .into_iter()
281            .filter(|field_id| all_vectors.get(field_id).is_none_or(|v| v.is_empty()))
282            .collect();
283        if !starved_built.is_empty() {
284            return Err(Error::Schema(format!(
285                "cannot retrain vector index: no training vectors could be collected \
286                 for built field(s) {starved_built:?}; the existing trained artifacts \
287                 are left in place"
288            )));
289        }
290
291        // Reset metadata and the ArcSwap set together. Old fixed-name artifact
292        // files are left in place until the atomic writer replaces them; this
293        // avoids a cancellation window and does not accumulate generations.
294        self.segment_manager
295            .update_vector_metadata_and_publish(&artifact_update, |meta| {
296                for field_id in &field_ids {
297                    if let Some(field_meta) = meta.vector_fields.get_mut(field_id) {
298                        field_meta.state = super::VectorIndexState::Flat;
299                        field_meta.centroids_file = None;
300                        field_meta.codebook_file = None;
301                    }
302                }
303                meta.refresh_total_vectors();
304            })
305            .await?;
306
307        log::info!("Reset vector index state to Flat, retraining from collected samples...");
308
309        self.train_and_publish_fields(
310            &dense_fields,
311            &all_vectors,
312            &total_vectors,
313            &artifact_update,
314        )
315        .await
316    }
317
318    // ========================================================================
319    // Helper methods
320    // ========================================================================
321
322    async fn reject_rebuild_with_ann_segments(
323        &self,
324        segment_ids: &[String],
325        field_ids: &[u32],
326    ) -> Result<()> {
327        for id_str in segment_ids {
328            let segment_id = SegmentId::from_hex(id_str)
329                .ok_or_else(|| Error::Corruption(format!("Invalid segment ID: {id_str}")))?;
330            let reader = SegmentReader::open_with_cache_blocks(
331                self.directory.as_ref(),
332                segment_id,
333                Arc::clone(&self.schema),
334                self.config.term_cache_blocks,
335                self.config.store_cache_blocks,
336            )
337            .await?;
338            Self::reject_ann_in_reader(&reader, id_str, field_ids)?;
339        }
340        Ok(())
341    }
342
343    fn reject_ann_in_reader(reader: &SegmentReader, id_str: &str, field_ids: &[u32]) -> Result<()> {
344        for &field_id in field_ids {
345            if matches!(
346                reader.vector_indexes().get(&field_id),
347                Some(crate::segment::VectorIndex::IVF(_))
348                    | Some(crate::segment::VectorIndex::ScaNN(_))
349            ) {
350                return Err(Error::Schema(format!(
351                    "cannot retrain vector artifacts for field {field_id}: segment {id_str} \
352                     already contains an IVF/ScaNN index built with the current generation; \
353                     rebuild requires all committed segments for the field to be flat"
354                )));
355            }
356        }
357        Ok(())
358    }
359
360    /// Get all dense vector fields that need ANN indexes
361    fn get_dense_vector_fields(&self) -> Vec<(Field, DenseVectorConfig)> {
362        self.schema
363            .fields()
364            .filter_map(|(field, entry)| {
365                if entry.field_type == FieldType::DenseVector && entry.indexed {
366                    entry
367                        .dense_vector_config
368                        .as_ref()
369                        // Only IVF-backed indexes require a global training
370                        // artifact. Standalone RaBitQ trains per segment; including
371                        // it here repeatedly sampled up to 100k vectors and then
372                        // returned without producing metadata.
373                        .filter(|c| c.uses_ivf())
374                        .map(|c| (field, c.clone()))
375                } else {
376                    None
377                }
378            })
379            .collect()
380    }
381
382    /// Get fields that need building (not already built)
383    async fn get_fields_to_build(
384        &self,
385        dense_fields: &[(Field, DenseVectorConfig)],
386    ) -> Vec<(Field, DenseVectorConfig)> {
387        let field_ids: Vec<u32> = dense_fields.iter().map(|(f, _)| f.0).collect();
388        let built: Vec<u32> = self
389            .segment_manager
390            .read_metadata(|meta| {
391                field_ids
392                    .iter()
393                    .filter(|fid| meta.is_field_built(**fid))
394                    .copied()
395                    .collect()
396            })
397            .await;
398        dense_fields
399            .iter()
400            .filter(|(field, _)| !built.contains(&field.0))
401            .cloned()
402            .collect()
403    }
404
405    /// Collect vectors from segments for training, with sampling for large datasets.
406    ///
407    /// K-means clustering converges well with ~100K samples, so we cap collection
408    /// per field to avoid loading millions of vectors into memory.
409    async fn collect_vectors_for_training(
410        &self,
411        segment_ids: &[String],
412        fields_to_build: &[(Field, DenseVectorConfig)],
413    ) -> Result<(FxHashMap<u32, Vec<Vec<f32>>>, FxHashMap<u32, usize>)> {
414        /// Maximum vectors per field for training. K-means converges well with ~100K samples.
415        const MAX_TRAINING_VECTORS: usize = 100_000;
416
417        let mut all_vectors: FxHashMap<u32, Vec<Vec<f32>>> = FxHashMap::default();
418        let mut total_vectors: FxHashMap<u32, usize> = FxHashMap::default();
419        let mut total_skipped = 0usize;
420        let field_ids: Vec<u32> = fields_to_build.iter().map(|(field, _)| field.0).collect();
421
422        for id_str in segment_ids {
423            let segment_id = SegmentId::from_hex(id_str)
424                .ok_or_else(|| Error::Corruption(format!("Invalid segment ID: {}", id_str)))?;
425            let reader = SegmentReader::open_with_cache_blocks(
426                self.directory.as_ref(),
427                segment_id,
428                Arc::clone(&self.schema),
429                self.config.term_cache_blocks,
430                self.config.store_cache_blocks,
431            )
432            .await?;
433
434            // `build_vector_index` is also effectively a retrain whenever
435            // metadata says Flat. A crash-interrupted rebuild from an older
436            // Hermes version can leave that state beside committed ANN
437            // segments, so validate generation safety during the same segment
438            // scan that collects the samples.
439            Self::reject_ann_in_reader(&reader, id_str, &field_ids)?;
440
441            for (field_id, lazy_flat) in reader.flat_vectors() {
442                if !fields_to_build.iter().any(|(f, _)| f.0 == *field_id) {
443                    continue;
444                }
445                let total = total_vectors.entry(*field_id).or_default();
446                *total = total.saturating_add(lazy_flat.num_vectors);
447                let entry = all_vectors.entry(*field_id).or_default();
448                let remaining = MAX_TRAINING_VECTORS.saturating_sub(entry.len());
449
450                if remaining == 0 {
451                    total_skipped += lazy_flat.num_vectors;
452                    continue;
453                }
454
455                let n = lazy_flat.num_vectors;
456                let dim = lazy_flat.dim;
457                let quant = lazy_flat.quantization;
458
459                // Determine which vector indices to collect
460                let indices: Vec<usize> = if n <= remaining {
461                    (0..n).collect()
462                } else {
463                    let step = (n / remaining).max(1);
464                    (0..n).step_by(step).take(remaining).collect()
465                };
466
467                if indices.len() < n {
468                    total_skipped += n - indices.len();
469                }
470
471                // Batch-read and dequantize instead of one-by-one get_vector().
472                // Read failures propagate: silently skipping vectors would
473                // train on an arbitrarily biased sample (or none at all) with
474                // no observability.
475                const BATCH: usize = 1024;
476                let mut f32_buf = vec![0f32; BATCH * dim];
477                for chunk in indices.chunks(BATCH) {
478                    // For contiguous ranges, use batch read
479                    let start = chunk[0];
480                    let end = *chunk.last().unwrap();
481                    if end - start + 1 == chunk.len() {
482                        // Contiguous — single batch read
483                        let batch_bytes = lazy_flat
484                            .read_vectors_batch(start, chunk.len())
485                            .await
486                            .map_err(crate::Error::Io)?;
487                        let floats = chunk.len() * dim;
488                        f32_buf.resize(floats, 0.0);
489                        crate::segment::dequantize_raw(
490                            batch_bytes.as_slice(),
491                            quant,
492                            floats,
493                            &mut f32_buf,
494                        )
495                        .map_err(crate::Error::Io)?;
496                        for i in 0..chunk.len() {
497                            entry.push(f32_buf[i * dim..(i + 1) * dim].to_vec());
498                        }
499                    } else {
500                        // Non-contiguous (sampled) — read individually but reuse buffer
501                        f32_buf.resize(dim, 0.0);
502                        for &idx in chunk {
503                            lazy_flat
504                                .read_vector_into(idx, &mut f32_buf)
505                                .await
506                                .map_err(crate::Error::Io)?;
507                            entry.push(f32_buf[..dim].to_vec());
508                        }
509                    }
510                }
511            }
512        }
513
514        if total_skipped > 0 {
515            let collected: usize = all_vectors.values().map(|v| v.len()).sum();
516            log::info!(
517                "Sampled {} vectors for training (skipped {}, max {} per field)",
518                collected,
519                total_skipped,
520                MAX_TRAINING_VECTORS,
521            );
522        }
523
524        Ok((all_vectors, total_vectors))
525    }
526
527    /// Train index for a single field
528    async fn train_field_index(
529        &self,
530        field: Field,
531        config: &DenseVectorConfig,
532        all_vectors: &FxHashMap<u32, Vec<Vec<f32>>>,
533        total_vectors: &FxHashMap<u32, usize>,
534    ) -> Result<Option<TrainedFieldUpdate>> {
535        let field_id = field.0;
536        let vectors = match all_vectors.get(&field_id) {
537            Some(v) if !v.is_empty() => v,
538            _ => return Ok(None),
539        };
540
541        let dim = config.dim;
542        let sample_count = vectors.len();
543        let corpus_count = total_vectors
544            .get(&field_id)
545            .copied()
546            .unwrap_or(sample_count);
547        // RaBitQ is trained independently per segment and does not need an
548        // index-level centroid artifact.
549        if !matches!(
550            config.index_type,
551            VectorIndexType::IvfRaBitQ | VectorIndexType::ScaNN
552        ) {
553            return Ok(None);
554        }
555
556        let num_clusters = effective_ivf_num_clusters(config, corpus_count, sample_count)?;
557
558        log::info!(
559            "Training vector index for field {} with {} sampled / {} total vectors, {} clusters (dim={})",
560            field_id,
561            sample_count,
562            corpus_count,
563            num_clusters,
564            dim,
565        );
566
567        let centroids_filename = format!("field_{}_centroids.bin", field_id);
568        let mut codebook_filename: Option<String> = None;
569
570        let actual_num_clusters = match config.index_type {
571            VectorIndexType::IvfRaBitQ => {
572                self.train_ivf_rabitq(
573                    field_id,
574                    dim,
575                    num_clusters,
576                    config.soar.clone(),
577                    vectors,
578                    &centroids_filename,
579                )
580                .await?
581            }
582            VectorIndexType::ScaNN => {
583                codebook_filename = Some(format!("field_{}_codebook.bin", field_id));
584                self.train_scann(
585                    field_id,
586                    dim,
587                    num_clusters,
588                    config.soar.clone(),
589                    vectors,
590                    &centroids_filename,
591                    codebook_filename.as_ref().unwrap(),
592                )
593                .await?
594            }
595            _ => unreachable!("non-IVF vector index returned above"),
596        };
597
598        Ok(Some(TrainedFieldUpdate {
599            field_id,
600            index_type: config.index_type,
601            vector_count: corpus_count,
602            num_clusters: actual_num_clusters,
603            centroids_file: centroids_filename,
604            codebook_file: codebook_filename,
605        }))
606    }
607
608    /// Serialize a trained structure to bincode and save to an index-level file.
609    async fn save_trained_artifact(
610        &self,
611        artifact: &impl serde::Serialize,
612        filename: &str,
613    ) -> Result<()> {
614        let temp_filename = format!("{filename}.tmp");
615        let temp_path = std::path::Path::new(&temp_filename);
616        let final_path = std::path::Path::new(filename);
617        let mut writer = self.directory.streaming_writer(temp_path).await?;
618        let encode_result = {
619            let mut limited = SizeLimitedWriter::new(
620                writer.as_mut(),
621                super::metadata::MAX_TRAINED_ARTIFACT_BYTES,
622            );
623            bincode::serde::encode_into_std_write(
624                artifact,
625                &mut limited,
626                bincode::config::standard(),
627            )
628        };
629        if let Err(error) = encode_result {
630            drop(writer);
631            let _ = self.directory.delete(temp_path).await;
632            return Err(Error::Serialization(format!(
633                "failed to serialize trained artifact '{filename}': {error}"
634            )));
635        }
636        if let Err(error) = writer.finish() {
637            let _ = self.directory.delete(temp_path).await;
638            return Err(Error::Io(error));
639        }
640        if let Err(error) = self.directory.rename(temp_path, final_path).await {
641            let _ = self.directory.delete(temp_path).await;
642            return Err(Error::Io(error));
643        }
644        self.directory.sync().await?;
645        Ok(())
646    }
647
648    /// Train IVF-RaBitQ centroids
649    async fn train_ivf_rabitq(
650        &self,
651        field_id: u32,
652        dim: usize,
653        num_clusters: usize,
654        soar: Option<crate::structures::SoarConfig>,
655        vectors: &[Vec<f32>],
656        centroids_filename: &str,
657    ) -> Result<usize> {
658        let mut coarse_config = crate::structures::CoarseConfig::new(dim, num_clusters);
659        if let Some(soar) = soar {
660            coarse_config = coarse_config.with_soar(soar);
661        }
662        let centroids = crate::structures::CoarseCentroids::train(&coarse_config, vectors);
663        self.save_trained_artifact(&centroids, centroids_filename)
664            .await?;
665
666        log::info!(
667            "Saved IVF-RaBitQ centroids for field {} ({} clusters, soar={})",
668            field_id,
669            centroids.num_clusters,
670            centroids.soar_config.is_some()
671        );
672        Ok(centroids.num_clusters as usize)
673    }
674
675    /// Train ScaNN (IVF-PQ) centroids and codebook
676    #[allow(clippy::too_many_arguments)]
677    async fn train_scann(
678        &self,
679        field_id: u32,
680        dim: usize,
681        num_clusters: usize,
682        soar: Option<crate::structures::SoarConfig>,
683        vectors: &[Vec<f32>],
684        centroids_filename: &str,
685        codebook_filename: &str,
686    ) -> Result<usize> {
687        let mut coarse_config = crate::structures::CoarseConfig::new(dim, num_clusters);
688        if let Some(soar) = soar {
689            coarse_config = coarse_config.with_soar(soar);
690        }
691        let centroids = crate::structures::CoarseCentroids::train(&coarse_config, vectors);
692        self.save_trained_artifact(&centroids, centroids_filename)
693            .await?;
694
695        let pq_config = crate::structures::PQConfig::new(dim);
696        let codebook = crate::structures::PQCodebook::train(pq_config, vectors, 10);
697        self.save_trained_artifact(&codebook, codebook_filename)
698            .await?;
699
700        log::info!(
701            "Saved ScaNN centroids and codebook for field {} ({} clusters)",
702            field_id,
703            centroids.num_clusters
704        );
705        Ok(centroids.num_clusters as usize)
706    }
707}
708
709#[cfg(test)]
710mod tests {
711    use super::*;
712
713    fn ivf_config(num_clusters: Option<usize>) -> DenseVectorConfig {
714        DenseVectorConfig::with_ivf(8, num_clusters, 4)
715    }
716
717    #[test]
718    fn effective_clusters_follow_corpus_heuristic_but_fit_sample() {
719        let config = ivf_config(None);
720
721        assert_eq!(
722            effective_ivf_num_clusters(&config, 1_000_000, 73).unwrap(),
723            73
724        );
725        assert_eq!(
726            effective_ivf_num_clusters(&config, 10_000, 1_000).unwrap(),
727            100
728        );
729    }
730
731    #[test]
732    fn effective_clusters_clamp_explicit_value_to_sample() {
733        let config = ivf_config(Some(256));
734        assert_eq!(
735            effective_ivf_num_clusters(&config, 1_000_000, 17).unwrap(),
736            17
737        );
738    }
739
740    #[test]
741    fn effective_clusters_reject_invalid_explicit_bounds() {
742        let zero = effective_ivf_num_clusters(&ivf_config(Some(0)), 10_000, 100)
743            .unwrap_err()
744            .to_string();
745        assert!(zero.contains("at least 1"));
746
747        let too_many =
748            effective_ivf_num_clusters(&ivf_config(Some(MAX_IVF_CLUSTERS + 1)), 10_000, 100)
749                .unwrap_err()
750                .to_string();
751        assert!(too_many.contains("must not exceed 4096"));
752    }
753
754    #[test]
755    fn effective_clusters_reject_empty_training_sample() {
756        let error = effective_ivf_num_clusters(&ivf_config(None), 10_000, 0)
757            .unwrap_err()
758            .to_string();
759        assert!(error.contains("without sample vectors"));
760    }
761
762    #[test]
763    fn artifact_writer_enforces_limit_without_writing_past_it() {
764        let mut output = Vec::new();
765        let mut writer = SizeLimitedWriter::new(&mut output, 3);
766        writer.write_all(&[1, 2]).unwrap();
767        let error = writer.write_all(&[3, 4]).unwrap_err().to_string();
768        assert!(error.contains("3-byte safety limit"), "{error}");
769        assert_eq!(output, vec![1, 2]);
770    }
771
772    // ===== rebuild destructive-downgrade regression tests =====
773
774    use std::path::Path;
775    use std::sync::atomic::{AtomicBool, Ordering};
776
777    use crate::directories::{
778        Directory, DirectoryWriter as DirectoryWriterTrait, FileHandle, RamDirectory, RangeReadFn,
779    };
780    use crate::dsl::{Document, SchemaBuilder};
781    use crate::index::{IndexConfig, IndexWriter};
782
783    const READ_FAIL_DOCS: usize = 5;
784    const READ_FAIL_DIM: usize = 4;
785    /// Flat entry layout of a single-field, flat-only `.vectors` file written
786    /// by the segment builder (data-first format): header (16 bytes) + raw f32
787    /// vectors + doc-id map + TOC + footer. Only the raw vector region is read
788    /// by training collection; segment open touches the header, doc-id map,
789    /// TOC, and footer, which all live outside this byte range.
790    const VEC_REGION_START: u64 = 16;
791    const VEC_REGION_END: u64 = VEC_REGION_START + (READ_FAIL_DOCS * READ_FAIL_DIM * 4) as u64;
792
793    /// RamDirectory wrapper whose `.vectors` handles fail range reads of the
794    /// raw vector region while `fail_vector_reads` is armed. Segment open
795    /// keeps succeeding, so exactly the training-collection batch reads fail —
796    /// the I/O the rebuild path used to swallow with `if let Ok`.
797    #[derive(Clone, Default)]
798    struct VectorReadFailDirectory {
799        inner: RamDirectory,
800        fail_vector_reads: Arc<AtomicBool>,
801    }
802
803    #[async_trait::async_trait]
804    impl Directory for VectorReadFailDirectory {
805        async fn exists(&self, path: &Path) -> std::io::Result<bool> {
806            self.inner.exists(path).await
807        }
808
809        async fn file_size(&self, path: &Path) -> std::io::Result<u64> {
810            self.inner.file_size(path).await
811        }
812
813        async fn open_read(&self, path: &Path) -> std::io::Result<FileHandle> {
814            self.inner.open_read(path).await
815        }
816
817        async fn read_range(
818            &self,
819            path: &Path,
820            range: std::ops::Range<u64>,
821        ) -> std::io::Result<crate::directories::OwnedBytes> {
822            self.inner.read_range(path, range).await
823        }
824
825        async fn list_files(&self, prefix: &Path) -> std::io::Result<Vec<std::path::PathBuf>> {
826            self.inner.list_files(prefix).await
827        }
828
829        async fn open_lazy(&self, path: &Path) -> std::io::Result<FileHandle> {
830            let handle = self.inner.open_lazy(path).await?;
831            if path.extension().is_some_and(|ext| ext == "vectors") {
832                let armed = Arc::clone(&self.fail_vector_reads);
833                let len = handle.len();
834                let read_fn: RangeReadFn = Arc::new(move |range: std::ops::Range<u64>| {
835                    let handle = handle.clone();
836                    let armed = Arc::clone(&armed);
837                    Box::pin(async move {
838                        if armed.load(Ordering::SeqCst)
839                            && range.start >= VEC_REGION_START
840                            && range.end <= VEC_REGION_END
841                        {
842                            return Err(std::io::Error::other("injected vector data read failure"));
843                        }
844                        handle.read_bytes_range(range).await
845                    })
846                });
847                return Ok(FileHandle::lazy(len, read_fn));
848            }
849            Ok(handle)
850        }
851    }
852
853    #[async_trait::async_trait]
854    impl DirectoryWriterTrait for VectorReadFailDirectory {
855        async fn write(&self, path: &Path, data: &[u8]) -> std::io::Result<()> {
856            self.inner.write(path, data).await
857        }
858
859        async fn delete(&self, path: &Path) -> std::io::Result<()> {
860            self.inner.delete(path).await
861        }
862
863        async fn rename(&self, from: &Path, to: &Path) -> std::io::Result<()> {
864            self.inner.rename(from, to).await
865        }
866
867        async fn sync(&self) -> std::io::Result<()> {
868            self.inner.sync().await
869        }
870
871        async fn streaming_writer(
872            &self,
873            path: &Path,
874        ) -> std::io::Result<Box<dyn crate::directories::StreamingWriter>> {
875            self.inner.streaming_writer(path).await
876        }
877    }
878
879    /// Regression: rebuild_vector_index used to durably reset Built fields to
880    /// Flat first and then swallow per-batch vector read errors with
881    /// `if let Ok` during training collection, reporting success while the
882    /// published trained generation had been destroyed. Read failures must
883    /// propagate, and the durable Built -> Flat reset must not happen.
884    #[tokio::test]
885    async fn rebuild_propagates_vector_read_errors_without_downgrading_built_state() {
886        let mut sb = SchemaBuilder::default();
887        let embedding = sb.add_dense_vector_field_with_config(
888            "embedding",
889            true,
890            true,
891            DenseVectorConfig::with_ivf(READ_FAIL_DIM, Some(1), 1),
892        );
893        let schema = sb.build();
894
895        let dir = VectorReadFailDirectory::default();
896        let config = IndexConfig {
897            merge_policy: Box::new(crate::merge::NoMergePolicy),
898            num_indexing_threads: 1,
899            ..Default::default()
900        };
901        let mut writer = IndexWriter::create(dir.clone(), schema, config)
902            .await
903            .unwrap();
904        for i in 0..READ_FAIL_DOCS {
905            let mut doc = Document::new();
906            doc.add_dense_vector(embedding, vec![i as f32 + 1.0; READ_FAIL_DIM]);
907            writer.add_document(doc).unwrap();
908        }
909        writer.commit().await.unwrap();
910        writer.build_vector_index().await.unwrap();
911        assert!(
912            writer
913                .segment_manager
914                .read_metadata(|meta| meta.is_field_built(embedding.0))
915                .await
916        );
917        assert!(writer.segment_manager.trained().is_some());
918
919        // Vector data reads now fail (transient I/O error).
920        dir.fail_vector_reads.store(true, Ordering::SeqCst);
921        let error = writer
922            .rebuild_vector_index()
923            .await
924            .expect_err("failed sample collection must fail the rebuild")
925            .to_string();
926        assert!(
927            error.contains("injected vector data read failure"),
928            "{error}"
929        );
930
931        // The published generation survives: no durable Built -> Flat reset.
932        assert!(
933            writer
934                .segment_manager
935                .read_metadata(|meta| meta.is_field_built(embedding.0))
936                .await,
937            "a failed rebuild must not durably downgrade the field to Flat"
938        );
939        assert!(
940            writer.segment_manager.trained().is_some(),
941            "a failed rebuild must not clear the published trained artifacts"
942        );
943    }
944
945    /// Regression: rebuild_vector_index used to return Ok(()) after durably
946    /// resetting a Built field to Flat even when no training vectors could be
947    /// collected at all, silently discarding the trained generation. An empty
948    /// training sample for a Built field must be a hard error raised BEFORE
949    /// the durable reset.
950    #[tokio::test]
951    async fn rebuild_errors_before_reset_when_built_field_has_no_training_vectors() {
952        let mut sb = SchemaBuilder::default();
953        let title = sb.add_text_field("title", true, true);
954        let embedding = sb.add_dense_vector_field_with_config(
955            "embedding",
956            true,
957            true,
958            DenseVectorConfig::with_ivf(4, Some(1), 1),
959        );
960        let schema = sb.build();
961
962        let dir = RamDirectory::new();
963        let config = IndexConfig {
964            merge_policy: Box::new(crate::merge::NoMergePolicy),
965            num_indexing_threads: 1,
966            ..Default::default()
967        };
968        let mut writer = IndexWriter::create(dir.clone(), schema, config)
969            .await
970            .unwrap();
971        // Committed segments carry no vectors for the field.
972        for i in 0..3 {
973            let mut doc = Document::new();
974            doc.add_text(title, format!("doc {i}"));
975            writer.add_document(doc).unwrap();
976        }
977        writer.commit().await.unwrap();
978
979        // Metadata says Built while no committed segment holds vectors for the
980        // field — the state a crash/degradation can leave behind. Rebuilding
981        // must refuse to destroy the referenced artifacts.
982        writer
983            .segment_manager
984            .update_metadata(|meta| {
985                meta.init_field(embedding.0, VectorIndexType::IvfRaBitQ);
986                meta.mark_field_built(
987                    embedding.0,
988                    5,
989                    1,
990                    format!("field_{}_centroids.bin", embedding.0),
991                    None,
992                );
993            })
994            .await
995            .unwrap();
996
997        let error = writer
998            .rebuild_vector_index()
999            .await
1000            .expect_err("an empty training sample must fail the rebuild")
1001            .to_string();
1002        assert!(error.contains("no training vectors"), "{error}");
1003        assert!(
1004            writer
1005                .segment_manager
1006                .read_metadata(|meta| meta.is_field_built(embedding.0))
1007                .await,
1008            "an empty training sample must not durably downgrade the field to Flat"
1009        );
1010    }
1011}