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//! `build_vector_index()` trains missing coarse-centroid generations;
5//! `retrain_vector_index()` replaces them. Both finish every committed ANN
6//! segment. Leaf codecs (TurboQuant) are derived, never trained.
7
8use std::io::Write;
9use std::sync::Arc;
10
11use rustc_hash::FxHashMap;
12
13use crate::directories::DirectoryWriter;
14use crate::dsl::{
15    BinaryDenseVectorConfig, BinaryIndexType, DenseVectorConfig, Field, FieldType, VectorIndexType,
16};
17use crate::error::{Error, Result};
18use crate::segment::{SegmentFiles, SegmentId, SegmentMeta};
19
20use super::IndexWriter;
21
22/// Maximum supported IVF centroid count. Query-side `nprobe` and serialized
23/// cluster identifiers use the same practical bound.
24const MAX_IVF_CLUSTERS: usize = 1_048_576;
25/// Faiss-style clustering quality floor: fewer points per centroid generally
26/// overfits the training sample and leaves unstable/empty cells.
27const MIN_TRAINING_POINTS_PER_CENTROID: usize = 39;
28/// Faiss-style clustering ceiling: more points per centroid multiply Lloyd
29/// cost without materially improving the codebook.
30const COARSE_TRAINING_POINTS_PER_CENTROID: usize = 256;
31/// Bound transient I/O/dequantization buffers independently of the configured
32/// total training sample budget.
33const MAX_SAMPLE_READ_BYTES: usize = 64 * 1024 * 1024;
34const SAMPLE_BLOCK: usize = 256;
35/// Generation-qualified filenames make retraining crash-safe: the currently
36/// published metadata never points at a file being overwritten in place.
37const VECTOR_ARTIFACT_PREFIX: &str = "vector_artifact_";
38
39struct TrainedFieldUpdate {
40    field_id: u32,
41    index_type: super::metadata::VectorFieldIndexType,
42    vector_count: usize,
43    num_clusters: usize,
44    centroids_file: String,
45    codebook_file: Option<String>,
46}
47
48enum TrainedFieldArtifacts {
49    /// IVF-TQ: only the coarse router is trained; the TQ leaf codec is
50    /// derived from the dimension.
51    FloatCentroids(crate::structures::CoarseCentroids),
52    Binary(crate::structures::BinaryCoarseQuantizer),
53}
54
55struct TrainedFieldModel {
56    update: TrainedFieldUpdate,
57    artifacts: TrainedFieldArtifacts,
58}
59
60#[derive(Clone)]
61enum IvfFieldConfig {
62    Float(DenseVectorConfig),
63    Binary(BinaryDenseVectorConfig),
64}
65
66impl IvfFieldConfig {
67    fn dim(&self) -> usize {
68        match self {
69            Self::Float(config) => config.dim,
70            Self::Binary(config) => config.dim,
71        }
72    }
73
74    fn index_type(&self) -> super::metadata::VectorFieldIndexType {
75        match self {
76            Self::Float(config) => config.index_type.into(),
77            Self::Binary(config) => config.index_type.into(),
78        }
79    }
80
81    fn num_clusters(&self) -> Option<usize> {
82        match self {
83            Self::Float(config) => config.num_clusters,
84            Self::Binary(config) => config.num_clusters,
85        }
86    }
87
88    fn optimal_num_clusters(&self, vector_count: usize) -> usize {
89        match self {
90            Self::Float(config) => config.optimal_num_clusters(vector_count),
91            Self::Binary(config) => config.optimal_num_clusters(vector_count),
92        }
93    }
94}
95
96enum TrainingSample {
97    Float(Vec<Vec<f32>>),
98    Binary(Vec<u8>),
99}
100
101#[derive(Clone, Copy, Debug, Eq, PartialEq)]
102enum VectorGenerationMode {
103    BuildMissing,
104    RetrainAll,
105}
106
107impl TrainingSample {
108    fn len(&self, dim: usize) -> usize {
109        match self {
110            Self::Float(vectors) => vectors.len(),
111            Self::Binary(codes) => codes.len() / dim.div_ceil(8),
112        }
113    }
114}
115
116/// Write adapter that rejects an artifact before its serialized form exceeds
117/// the same bound enforced by the loader. Encoding directly through this
118/// adapter avoids materializing a second, potentially hundreds-of-megabytes
119/// copy of the trained structure.
120struct SizeLimitedWriter<'a, W: Write + ?Sized> {
121    inner: &'a mut W,
122    written: usize,
123    limit: usize,
124}
125
126impl<'a, W: Write + ?Sized> SizeLimitedWriter<'a, W> {
127    fn new(inner: &'a mut W, limit: usize) -> Self {
128        Self {
129            inner,
130            written: 0,
131            limit,
132        }
133    }
134}
135
136impl<W: Write + ?Sized> Write for SizeLimitedWriter<'_, W> {
137    fn write(&mut self, buffer: &[u8]) -> std::io::Result<usize> {
138        let next_size = self
139            .written
140            .checked_add(buffer.len())
141            .ok_or_else(|| std::io::Error::other("trained artifact size overflow"))?;
142        if next_size > self.limit {
143            return Err(std::io::Error::new(
144                std::io::ErrorKind::InvalidData,
145                format!(
146                    "trained artifact exceeds the {}-byte safety limit",
147                    self.limit
148                ),
149            ));
150        }
151        let written = self.inner.write(buffer)?;
152        self.written += written;
153        Ok(written)
154    }
155
156    fn flush(&mut self) -> std::io::Result<()> {
157        self.inner.flush()
158    }
159}
160
161fn validate_explicit_cluster_count(num_clusters: Option<usize>) -> Result<()> {
162    match num_clusters {
163        Some(0) => Err(Error::Schema(
164            "dense vector num_clusters must be at least 1".to_string(),
165        )),
166        Some(value) if value > MAX_IVF_CLUSTERS => Err(Error::Schema(format!(
167            "dense vector num_clusters must not exceed {MAX_IVF_CLUSTERS}, got {value}"
168        ))),
169        _ => Ok(()),
170    }
171}
172
173fn effective_field_num_clusters(
174    config: &IvfFieldConfig,
175    corpus_count: usize,
176    sample_count: usize,
177) -> Result<usize> {
178    if sample_count == 0 {
179        return Err(Error::Schema(
180            "cannot train an IVF vector index without sample vectors".to_string(),
181        ));
182    }
183    validate_explicit_cluster_count(config.num_clusters())?;
184    let centroid_bytes = match config {
185        IvfFieldConfig::Float(config) => config.dim.saturating_mul(size_of::<f32>()),
186        IvfFieldConfig::Binary(config) => config.dim.div_ceil(8),
187    };
188    let artifact_limit = super::metadata::MAX_TRAINED_ARTIFACT_BYTES
189        .saturating_sub(1024)
190        .checked_div(centroid_bytes.max(1))
191        .unwrap_or(0)
192        .max(1);
193    let quality_limit = if config.num_clusters().is_some() {
194        sample_count
195    } else {
196        (sample_count / MIN_TRAINING_POINTS_PER_CENTROID)
197            .max(16)
198            .min(sample_count)
199    };
200    let requested = config.optimal_num_clusters(corpus_count);
201    if config.num_clusters().is_some() && requested > artifact_limit {
202        return Err(Error::Schema(format!(
203            "configured IVF codebook needs {} bytes for {} centroids, exceeding the {}-byte artifact limit",
204            requested.saturating_mul(centroid_bytes),
205            requested,
206            super::metadata::MAX_TRAINED_ARTIFACT_BYTES,
207        )));
208    }
209    Ok(requested.min(quality_limit).min(artifact_limit))
210}
211
212fn training_sample_limit(
213    max_samples: usize,
214    max_bytes: usize,
215    bytes_per_sample: usize,
216) -> Result<usize> {
217    if max_samples == 0 || max_bytes == 0 || bytes_per_sample == 0 {
218        return Err(Error::Schema(
219            "vector training sample count, memory budget, and vector size must be greater than zero"
220                .into(),
221        ));
222    }
223    let memory_limited = max_bytes / bytes_per_sample;
224    if memory_limited == 0 {
225        return Err(Error::Schema(format!(
226            "vector training memory budget ({max_bytes} bytes) cannot hold one {bytes_per_sample}-byte sample"
227        )));
228    }
229    Ok(max_samples.min(memory_limited))
230}
231
232/// Validate the configured centroid count and cap it to the training sample.
233///
234/// Corpus size drives the automatic heuristic, but training cannot produce
235/// more distinct centroids than the number of sampled vectors. Keeping this
236/// decision here avoids relying on a panic-prone, implicit clamp inside the
237/// trainer and gives callers a schema error for invalid explicit values.
238#[cfg(test)]
239fn effective_ivf_num_clusters(
240    config: &DenseVectorConfig,
241    corpus_count: usize,
242    sample_count: usize,
243) -> Result<usize> {
244    if sample_count == 0 {
245        return Err(Error::Schema(
246            "cannot train an IVF vector index without sample vectors".to_string(),
247        ));
248    }
249
250    effective_field_num_clusters(
251        &IvfFieldConfig::Float(config.clone()),
252        corpus_count,
253        sample_count,
254    )
255}
256
257impl<D: DirectoryWriter + 'static> IndexWriter<D> {
258    /// Train vector index from accumulated Flat vectors (manual, not auto-triggered).
259    ///
260    /// 1. Acquires a stable segment snapshot.
261    /// 2. Trains missing coarse-centroid generations.
262    /// 3. Stages ANN replacements for every affected segment.
263    /// 4. Publishes the complete segment/codebook generation atomically.
264    pub async fn build_vector_index(&self) -> Result<()> {
265        self.build_vector_generation(VectorGenerationMode::BuildMissing)
266            .await
267    }
268
269    /// Train a fresh global codebook from the current corpus and rebuild every
270    /// ANN segment into that generation. The replacement is atomic for search
271    /// readers: the old segment/codebook pair remains live until all new files
272    /// have been staged and durably committed together.
273    pub async fn retrain_vector_index(&self) -> Result<()> {
274        self.build_vector_generation(VectorGenerationMode::RetrainAll)
275            .await
276    }
277
278    async fn build_vector_generation(&self, mode: VectorGenerationMode) -> Result<()> {
279        let dense_fields = self.get_ivf_vector_fields();
280        if dense_fields.is_empty() {
281            log::info!("No dense vector fields configured for ANN indexing");
282            return Ok(());
283        }
284
285        let artifact_update = self.segment_manager.begin_vector_artifact_update().await?;
286        self.cleanup_unreferenced_vector_artifacts().await;
287
288        let fields_to_train = match mode {
289            VectorGenerationMode::BuildMissing => self.get_fields_to_build(&dense_fields).await,
290            VectorGenerationMode::RetrainAll => dense_fields.clone(),
291        };
292        for (_, config) in &fields_to_train {
293            validate_explicit_cluster_count(config.num_clusters())?;
294        }
295
296        let snapshot = self.segment_manager.acquire_snapshot().await;
297        if snapshot.is_empty() {
298            if mode == VectorGenerationMode::RetrainAll {
299                return Err(Error::Schema(
300                    "cannot retrain vector centroids without committed segments".into(),
301                ));
302            }
303            return Ok(());
304        }
305
306        let mut candidate_metadata = self.segment_manager.read_metadata(Clone::clone).await;
307        if !fields_to_train.is_empty() {
308            let total_vectors = self
309                .count_vectors_for_training(
310                    snapshot.segment_ids(),
311                    &fields_to_train,
312                    mode == VectorGenerationMode::BuildMissing,
313                )
314                .await?;
315            let artifact_generation = SegmentId::new().to_hex();
316            let updates = self
317                .train_fields(
318                    snapshot.segment_ids(),
319                    &fields_to_train,
320                    &total_vectors,
321                    &artifact_generation,
322                )
323                .await?;
324            for update in &updates {
325                candidate_metadata.init_field(update.field_id, update.index_type);
326                candidate_metadata.mark_field_built(
327                    update.field_id,
328                    update.vector_count,
329                    update.num_clusters,
330                    update.centroids_file.clone(),
331                    update.codebook_file.clone(),
332                );
333            }
334        }
335
336        let target_field_ids = dense_fields
337            .iter()
338            .filter_map(|(field, _)| {
339                candidate_metadata
340                    .is_field_built(field.0)
341                    .then_some(field.0)
342            })
343            .collect::<Vec<_>>();
344        if target_field_ids.is_empty() {
345            return Ok(());
346        }
347
348        let candidate_trained = super::IndexMetadata::try_load_trained_from_fields(
349            &candidate_metadata.vector_fields,
350            self.schema.as_ref(),
351            self.directory.as_ref(),
352        )
353        .await?
354        .map(Arc::new)
355        .ok_or_else(|| Error::Internal("candidate vector generation has no artifacts".into()))?;
356
357        let staged = self
358            .segment_manager
359            .stage_vector_generation(
360                &artifact_update,
361                snapshot.segment_ids(),
362                &target_field_ids,
363                Arc::clone(&candidate_trained),
364                mode == VectorGenerationMode::RetrainAll,
365            )
366            .await?;
367        self.segment_manager
368            .publish_vector_generation(
369                &artifact_update,
370                candidate_metadata.vector_fields,
371                candidate_trained,
372                staged,
373            )
374            .await?;
375
376        // Old readers retain the old snapshot and deserialized codebook. Once
377        // this local training snapshot drops, retired source files can be
378        // reclaimed. Reopening producers after the lease sees only the new set.
379        drop(snapshot);
380        drop(artifact_update);
381
382        // A producer that started while training was gated writes flat data.
383        // Catch already committed outputs; later commits carry their own
384        // targeted upgrade marker in PreparedSegment.
385        self.segment_manager
386            .rewrite_vector_segments(&target_field_ids)
387            .await?;
388        self.cleanup_unreferenced_vector_artifacts().await;
389        log::info!(
390            "Dense vector ANN generation {:?} complete for {} field(s)",
391            mode,
392            target_field_ids.len(),
393        );
394        Ok(())
395    }
396
397    async fn train_fields(
398        &self,
399        segment_ids: &[String],
400        fields: &[(Field, IvfFieldConfig)],
401        total_vectors: &FxHashMap<u32, usize>,
402        artifact_generation: &str,
403    ) -> Result<Vec<TrainedFieldUpdate>> {
404        let training_pool = self.segment_manager.background_cpu_pool();
405        let mut missing = Vec::new();
406        let mut updates = Vec::with_capacity(fields.len());
407        for (field, config) in fields {
408            // Sample collection and training are both field-serial. At most
409            // one bounded sample, one field's clustering scratch, and one
410            // generated artifact set can coexist.
411            let corpus_count = total_vectors.get(&field.0).copied().unwrap_or(0);
412            let Some(sample) = self
413                .collect_training_sample(segment_ids, *field, config, corpus_count)
414                .await?
415            else {
416                missing.push(field.0);
417                continue;
418            };
419            let model = crate::segment::block_in_place_if_multithread(|| {
420                training_pool.install(|| {
421                    Self::train_field_model(
422                        *field,
423                        config,
424                        &sample,
425                        corpus_count,
426                        artifact_generation,
427                    )
428                })
429            })?;
430            // Training artifacts own everything needed for persistence. Drop
431            // the potentially multi-gigabyte sample before async file I/O.
432            drop(sample);
433            updates.push(self.save_trained_field(model).await?);
434        }
435        if updates.is_empty() && !fields.is_empty() {
436            return Err(Error::Schema(format!(
437                "cannot train vector centroids: no committed vectors for field(s) {missing:?}"
438            )));
439        }
440        if !missing.is_empty() {
441            log::info!(
442                "Skipping dense vector field(s) {missing:?}: the current corpus contains no vectors"
443            );
444        }
445        Ok(updates)
446    }
447
448    /// Remove abandoned generation-qualified artifacts from cancelled or
449    /// crash-interrupted attempts. The metadata references are the complete
450    /// live set, and the exclusive update lease prevents another trainer from
451    /// creating a candidate concurrently with this sweep.
452    async fn cleanup_unreferenced_vector_artifacts(&self) {
453        let referenced = self
454            .segment_manager
455            .read_metadata(|metadata| {
456                metadata
457                    .vector_fields
458                    .values()
459                    .flat_map(|field| {
460                        field
461                            .centroids_file
462                            .iter()
463                            .chain(field.codebook_file.iter())
464                    })
465                    .cloned()
466                    .collect::<std::collections::HashSet<_>>()
467            })
468            .await;
469        let files = match self.directory.list_files(std::path::Path::new("")).await {
470            Ok(files) => files,
471            Err(error) => {
472                log::warn!("[trained] failed listing abandoned dense vector artifacts: {error}");
473                return;
474            }
475        };
476        for path in files {
477            let Some(name) = path.file_name().and_then(|name| name.to_str()) else {
478                continue;
479            };
480            if !name.starts_with(VECTOR_ARTIFACT_PREFIX)
481                || referenced.contains(path.to_string_lossy().as_ref())
482            {
483                continue;
484            }
485            if let Err(error) = self.directory.delete(&path).await
486                && error.kind() != std::io::ErrorKind::NotFound
487            {
488                log::warn!("[trained] failed deleting abandoned artifact {path:?}: {error}");
489            }
490        }
491    }
492
493    // ========================================================================
494    // Helper methods
495    // ========================================================================
496
497    fn reject_ann_fields(ann_fields: &[u32], id_str: &str, field_ids: &[u32]) -> Result<()> {
498        for &field_id in field_ids {
499            if ann_fields.binary_search(&field_id).is_ok() {
500                return Err(Error::Schema(format!(
501                    "metadata-flat field {field_id} already has ANN data in segment {id_str}; \
502                     recreate the index instead of mixing vector generations"
503                )));
504            }
505        }
506        Ok(())
507    }
508
509    /// Open only selected flat-vector fields plus the tiny segment metadata.
510    /// Training does not need term dictionaries, stores, sparse structures, or
511    /// corpus-sized ANN run columns, and must not pin those transient readers.
512    async fn load_training_vectors(
513        &self,
514        segment_id: SegmentId,
515        field_ids: &[u32],
516    ) -> Result<crate::segment::reader::loader::VectorsFileData> {
517        let files = SegmentFiles::new(segment_id.0);
518        let meta_bytes = self
519            .directory
520            .open_read(&files.meta)
521            .await?
522            .read_bytes()
523            .await?;
524        let meta = SegmentMeta::deserialize(meta_bytes.as_slice())?;
525        if meta.id != segment_id.0 {
526            return Err(Error::Corruption(format!(
527                "segment metadata ID {:032x} does not match file ID {}",
528                meta.id,
529                segment_id.to_hex(),
530            )));
531        }
532        crate::segment::reader::loader::load_flat_vectors_file(
533            self.directory.as_ref(),
534            &files,
535            self.schema.as_ref(),
536            meta.num_docs,
537            field_ids,
538        )
539        .await
540    }
541
542    /// Get all dense vector fields that need ANN indexes
543    fn get_ivf_vector_fields(&self) -> Vec<(Field, IvfFieldConfig)> {
544        self.schema
545            .fields()
546            .filter_map(|(field, entry)| {
547                if entry.field_type == FieldType::DenseVector && entry.indexed {
548                    entry
549                        .dense_vector_config
550                        .as_ref()
551                        // Flat is a pre-build storage state; the production ANN
552                        // path is trained once and shared by every segment.
553                        .filter(|c| c.uses_ivf())
554                        .map(|c| (field, IvfFieldConfig::Float(c.clone())))
555                } else if entry.field_type == FieldType::BinaryDenseVector && entry.indexed {
556                    entry
557                        .binary_dense_vector_config
558                        .as_ref()
559                        .filter(|config| config.index_type == BinaryIndexType::Ivf)
560                        .map(|config| (field, IvfFieldConfig::Binary(config.clone())))
561                } else {
562                    None
563                }
564            })
565            .collect()
566    }
567
568    /// Get fields that need building (not already built)
569    async fn get_fields_to_build(
570        &self,
571        dense_fields: &[(Field, IvfFieldConfig)],
572    ) -> Vec<(Field, IvfFieldConfig)> {
573        let field_ids: Vec<u32> = dense_fields.iter().map(|(f, _)| f.0).collect();
574        let built: Vec<u32> = self
575            .segment_manager
576            .read_metadata(|meta| {
577                field_ids
578                    .iter()
579                    .filter(|fid| meta.is_field_built(**fid))
580                    .copied()
581                    .collect()
582            })
583            .await;
584        dense_fields
585            .iter()
586            .filter(|(field, _)| !built.contains(&field.0))
587            .cloned()
588            .collect()
589    }
590
591    /// Count every configured field without reading any vector payload bytes.
592    async fn count_vectors_for_training(
593        &self,
594        segment_ids: &[String],
595        fields_to_build: &[(Field, IvfFieldConfig)],
596        require_flat_generation: bool,
597    ) -> Result<FxHashMap<u32, usize>> {
598        let mut total_vectors: FxHashMap<u32, usize> = FxHashMap::default();
599        let field_ids: Vec<u32> = fields_to_build.iter().map(|(field, _)| field.0).collect();
600
601        // Initial construction rejects
602        // ANN payloads for metadata-flat fields; an explicit retrain reads the
603        // exact flat vectors retained beside the current ANN generation.
604        for id_str in segment_ids {
605            let segment_id = SegmentId::from_hex(id_str)
606                .ok_or_else(|| Error::Corruption(format!("Invalid segment ID: {}", id_str)))?;
607            let vectors = self.load_training_vectors(segment_id, &field_ids).await?;
608
609            if require_flat_generation {
610                Self::reject_ann_fields(&vectors.ann_fields, id_str, &field_ids)?;
611            }
612
613            for (field, _) in fields_to_build {
614                if let Some(flat) = vectors.flat_vectors.get(&field.0) {
615                    let total = total_vectors.entry(field.0).or_default();
616                    *total = total.checked_add(flat.num_vectors).ok_or_else(|| {
617                        Error::Corruption(format!(
618                            "vector count overflows usize for field {}",
619                            field.0,
620                        ))
621                    })?;
622                }
623            }
624        }
625        Ok(total_vectors)
626    }
627
628    /// Fetch one deterministic, uniform field sample from the pinned segment
629    /// snapshot. Only selected ranges are read; all other corpus vectors stay
630    /// on disk. The caller trains and drops this sample before moving to the
631    /// next field.
632    async fn collect_training_sample(
633        &self,
634        segment_ids: &[String],
635        field: Field,
636        config: &IvfFieldConfig,
637        total: usize,
638    ) -> Result<Option<TrainingSample>> {
639        if total == 0 {
640            return Ok(None);
641        }
642        let bytes_per_sample = match config {
643            IvfFieldConfig::Float(config) => config
644                .dim
645                .checked_mul(size_of::<f32>())
646                .ok_or_else(|| Error::Schema("float training vector size overflows".into()))?,
647            IvfFieldConfig::Binary(config) => config.dim.div_ceil(8),
648        };
649        let limit = training_sample_limit(
650            self.config.vector_training_max_samples,
651            self.config.vector_training_memory_bytes,
652            bytes_per_sample,
653        )?;
654        let take = total.min(limit);
655        let mut rng = <rand::rngs::StdRng as rand::SeedableRng>::seed_from_u64(
656            0x4845_524d_4553_4956 ^ field.0 as u64 ^ total as u64,
657        );
658        let mut ordinals = Vec::with_capacity(take);
659        if take == total {
660            ordinals.extend(0..total);
661        } else {
662            let blocks = take.div_ceil(SAMPLE_BLOCK);
663            for block in 0..blocks {
664                let block_len = SAMPLE_BLOCK.min(take - ordinals.len());
665                let stratum_start = block.saturating_mul(total) / blocks;
666                let stratum_end = (block + 1).saturating_mul(total) / blocks;
667                let latest_start = stratum_end.saturating_sub(block_len);
668                let start = if latest_start > stratum_start {
669                    rand::Rng::random_range(&mut rng, stratum_start..=latest_start)
670                } else {
671                    stratum_start
672                };
673                ordinals.extend(start..start + block_len);
674            }
675        }
676
677        let mut sample = match config {
678            IvfFieldConfig::Float(_) => TrainingSample::Float(Vec::with_capacity(take)),
679            IvfFieldConfig::Binary(_) => TrainingSample::Binary(Vec::with_capacity(
680                take.checked_mul(bytes_per_sample)
681                    .ok_or_else(|| Error::Schema("binary training sample size overflows".into()))?,
682            )),
683        };
684        let max_read_vectors = (MAX_SAMPLE_READ_BYTES / bytes_per_sample.max(1)).max(1);
685        let mut global_offset = 0usize;
686        let mut cursor = 0usize;
687        let field_ids = [field.0];
688
689        for id_str in segment_ids {
690            let segment_id = SegmentId::from_hex(id_str)
691                .ok_or_else(|| Error::Corruption(format!("Invalid segment ID: {id_str}")))?;
692            let vectors = self.load_training_vectors(segment_id, &field_ids).await?;
693
694            let Some(lazy_flat) = vectors.flat_vectors.get(&field.0) else {
695                continue;
696            };
697            let base = global_offset;
698            let end = base.checked_add(lazy_flat.num_vectors).ok_or_else(|| {
699                Error::Corruption(format!("vector offset overflows for field {}", field.0))
700            })?;
701            global_offset = end;
702            let first = cursor;
703            while cursor < ordinals.len() && ordinals[cursor] < end {
704                cursor += 1;
705            }
706            let selected = &ordinals[first..cursor];
707            let mut run_start = 0;
708            while run_start < selected.len() {
709                let mut run_end = run_start + 1;
710                while run_end < selected.len()
711                    && run_end - run_start < max_read_vectors
712                    && selected[run_end] == selected[run_end - 1] + 1
713                {
714                    run_end += 1;
715                }
716                let local_start = selected[run_start] - base;
717                let run_len = run_end - run_start;
718                let bytes = lazy_flat
719                    .read_vectors_batch(local_start, run_len)
720                    .await
721                    .map_err(crate::Error::Io)?;
722                match &mut sample {
723                    TrainingSample::Binary(codes) => {
724                        let expected = run_len.checked_mul(bytes_per_sample).ok_or_else(|| {
725                            Error::Corruption("binary sample read size overflows".into())
726                        })?;
727                        if bytes.len() != expected {
728                            return Err(Error::Corruption(format!(
729                                "binary sample read returned {} bytes, expected {expected}",
730                                bytes.len(),
731                            )));
732                        }
733                        codes.extend_from_slice(bytes.as_slice());
734                    }
735                    TrainingSample::Float(vectors) => {
736                        let dim = lazy_flat.dim;
737                        let float_count = run_len.checked_mul(dim).ok_or_else(|| {
738                            Error::Corruption("float sample read size overflows".into())
739                        })?;
740                        let mut decoded = vec![0.0; float_count];
741                        crate::segment::dequantize_raw(
742                            bytes.as_slice(),
743                            lazy_flat.quantization,
744                            decoded.len(),
745                            &mut decoded,
746                        )
747                        .map_err(crate::Error::Io)?;
748                        vectors.extend(decoded.chunks_exact(dim).map(<[f32]>::to_vec));
749                    }
750                }
751                run_start = run_end;
752            }
753        }
754
755        let collected = sample.len(config.dim());
756        if global_offset != total || cursor != take || collected != take {
757            return Err(Error::Corruption(format!(
758                "training sample coverage mismatch for field {}: counted={total}, traversed={global_offset}, selected={cursor}, collected={collected}",
759                field.0,
760            )));
761        }
762        if collected < total {
763            log::info!(
764                "Sampled {} / {} dense vectors for field {} (max {} vectors / {} resident)",
765                collected,
766                total,
767                field.0,
768                self.config.vector_training_max_samples,
769                crate::format_bytes(self.config.vector_training_memory_bytes as u64),
770            );
771        }
772        Ok(Some(sample))
773    }
774
775    /// Train one field. Called from the shared bounded Rayon pool, so fields
776    /// and each field's internal clustering work compose without extra pools.
777    fn train_field_model(
778        field: Field,
779        config: &IvfFieldConfig,
780        sample: &TrainingSample,
781        corpus_count: usize,
782        artifact_generation: &str,
783    ) -> Result<TrainedFieldModel> {
784        let field_id = field.0;
785        let dim = config.dim();
786        let sample_count = sample.len(dim);
787        if sample_count == 0 || corpus_count == 0 {
788            return Err(Error::Internal(format!(
789                "empty training sample for non-empty field {field_id}"
790            )));
791        }
792        let num_clusters = effective_field_num_clusters(config, corpus_count, sample_count)?;
793
794        log::info!(
795            "Training dense vector index for field {} with {} sampled / {} total vectors, {} clusters (dim={})",
796            field_id,
797            sample_count,
798            corpus_count,
799            num_clusters,
800            dim,
801        );
802
803        let centroids_filename =
804            format!("{VECTOR_ARTIFACT_PREFIX}{artifact_generation}_field_{field_id}_centroids.bin");
805
806        let artifacts = match (config, sample) {
807            (IvfFieldConfig::Float(config), TrainingSample::Float(vectors))
808                if config.index_type == VectorIndexType::IvfTq =>
809            {
810                let mut coarse_config = crate::structures::CoarseConfig::new(dim, num_clusters)
811                    .with_routing(config.ivf_routing);
812                if let Some(soar) = config.soar.clone() {
813                    coarse_config = coarse_config.with_soar(soar);
814                }
815                // Faiss-style clustering ceiling: past ~256 points per
816                // centroid, extra samples multiply every Lloyd iteration
817                // without materially moving the centroids. Stride-subsample
818                // the (already stratified) training set past that.
819                let ceiling = num_clusters.saturating_mul(COARSE_TRAINING_POINTS_PER_CENTROID);
820                let training_set: Vec<Vec<f32>> = if vectors.len() > ceiling && ceiling > 0 {
821                    log::info!(
822                        "Field {field_id}: capping coarse training at {ceiling} of {} samples \
823                         ({COARSE_TRAINING_POINTS_PER_CENTROID}/centroid)",
824                        vectors.len(),
825                    );
826                    (0..ceiling)
827                        .map(|index| vectors[index.saturating_mul(vectors.len()) / ceiling].clone())
828                        .collect()
829                } else {
830                    Vec::new()
831                };
832                let training_ref: &[Vec<f32>] = if training_set.is_empty() {
833                    vectors
834                } else {
835                    &training_set
836                };
837                TrainedFieldArtifacts::FloatCentroids(crate::structures::CoarseCentroids::train(
838                    &coarse_config,
839                    training_ref,
840                ))
841            }
842            (IvfFieldConfig::Binary(config), TrainingSample::Binary(codes)) => {
843                let mut binary_config = crate::structures::BinaryIvfConfig::new(dim, num_clusters);
844                binary_config.max_train_samples = sample_count;
845                binary_config.routing = config.ivf_routing;
846                TrainedFieldArtifacts::Binary(
847                    crate::structures::BinaryCoarseQuantizer::train(
848                        binary_config,
849                        codes,
850                        sample_count,
851                    )
852                    .map_err(Error::Io)?,
853                )
854            }
855            _ => {
856                return Err(Error::Internal(format!(
857                    "training sample kind does not match field {field_id}"
858                )));
859            }
860        };
861
862        let actual_num_clusters = match &artifacts {
863            TrainedFieldArtifacts::FloatCentroids(centroids) => centroids.num_clusters as usize,
864            TrainedFieldArtifacts::Binary(quantizer) => quantizer.num_clusters as usize,
865        };
866        Ok(TrainedFieldModel {
867            update: TrainedFieldUpdate {
868                field_id,
869                index_type: config.index_type(),
870                vector_count: corpus_count,
871                num_clusters: actual_num_clusters,
872                centroids_file: centroids_filename,
873                codebook_file: None,
874            },
875            artifacts,
876        })
877    }
878
879    async fn save_trained_field(&self, model: TrainedFieldModel) -> Result<TrainedFieldUpdate> {
880        let TrainedFieldModel { update, artifacts } = model;
881        match artifacts {
882            TrainedFieldArtifacts::FloatCentroids(centroids) => {
883                self.save_trained_artifact(&centroids, &update.centroids_file)
884                    .await?;
885                log::info!(
886                    "Saved IVF-TQ coarse artifact for field {} ({} clusters; leaf codec is derived)",
887                    update.field_id,
888                    centroids.num_clusters,
889                );
890            }
891            TrainedFieldArtifacts::Binary(quantizer) => {
892                self.save_trained_artifact(&quantizer, &update.centroids_file)
893                    .await?;
894                log::info!(
895                    "Saved binary IVF artifact for field {} ({} clusters)",
896                    update.field_id,
897                    quantizer.num_clusters,
898                );
899            }
900        }
901        Ok(update)
902    }
903
904    /// Serialize a trained structure to bincode and save to an index-level file.
905    async fn save_trained_artifact(
906        &self,
907        artifact: &impl serde::Serialize,
908        filename: &str,
909    ) -> Result<()> {
910        let temp_filename = format!("{filename}.tmp");
911        let temp_path = std::path::Path::new(&temp_filename);
912        let final_path = std::path::Path::new(filename);
913        let mut writer = self.directory.streaming_writer(temp_path).await?;
914        let encode_result = {
915            let mut limited = SizeLimitedWriter::new(
916                writer.as_mut(),
917                super::metadata::MAX_TRAINED_ARTIFACT_BYTES,
918            );
919            bincode::serde::encode_into_std_write(
920                artifact,
921                &mut limited,
922                bincode::config::standard(),
923            )
924        };
925        if let Err(error) = encode_result {
926            drop(writer);
927            let _ = self.directory.delete(temp_path).await;
928            return Err(Error::Serialization(format!(
929                "failed to serialize trained artifact '{filename}': {error}"
930            )));
931        }
932        if let Err(error) = writer.finish() {
933            let _ = self.directory.delete(temp_path).await;
934            return Err(Error::Io(error));
935        }
936        if let Err(error) = self.directory.rename(temp_path, final_path).await {
937            let _ = self.directory.delete(temp_path).await;
938            return Err(Error::Io(error));
939        }
940        self.directory.sync().await?;
941        Ok(())
942    }
943}
944
945#[cfg(test)]
946mod tests {
947    use super::*;
948
949    fn ivf_config(num_clusters: Option<usize>) -> DenseVectorConfig {
950        DenseVectorConfig::ivf_tq(8, num_clusters, 4)
951    }
952
953    #[test]
954    fn effective_clusters_follow_corpus_heuristic_but_fit_sample() {
955        let config = ivf_config(None);
956
957        assert_eq!(
958            effective_ivf_num_clusters(&config, 1_000_000, 73).unwrap(),
959            16
960        );
961        assert_eq!(
962            effective_ivf_num_clusters(&config, 10_000, 1_000).unwrap(),
963            25
964        );
965    }
966
967    #[test]
968    fn effective_clusters_clamp_explicit_value_to_sample() {
969        let config = ivf_config(Some(256));
970        assert_eq!(
971            effective_ivf_num_clusters(&config, 1_000_000, 17).unwrap(),
972            17
973        );
974    }
975
976    #[test]
977    fn effective_clusters_reject_invalid_explicit_bounds() {
978        let zero = effective_ivf_num_clusters(&ivf_config(Some(0)), 10_000, 100)
979            .unwrap_err()
980            .to_string();
981        assert!(zero.contains("at least 1"));
982
983        let too_many =
984            effective_ivf_num_clusters(&ivf_config(Some(MAX_IVF_CLUSTERS + 1)), 10_000, 100)
985                .unwrap_err()
986                .to_string();
987        assert!(too_many.contains("must not exceed 1048576"));
988    }
989
990    #[test]
991    fn effective_clusters_reject_empty_training_sample() {
992        let error = effective_ivf_num_clusters(&ivf_config(None), 10_000, 0)
993            .unwrap_err()
994            .to_string();
995        assert!(error.contains("without sample vectors"));
996    }
997
998    #[test]
999    fn training_sample_limit_honors_both_cli_bounds() {
1000        assert_eq!(training_sample_limit(10_000_000, 4_096, 4).unwrap(), 1_024);
1001        assert_eq!(training_sample_limit(100, 4_096, 4).unwrap(), 100);
1002        let error = training_sample_limit(100, 3, 4).unwrap_err().to_string();
1003        assert!(error.contains("cannot hold one"), "{error}");
1004    }
1005
1006    #[test]
1007    fn artifact_writer_enforces_limit_without_writing_past_it() {
1008        let mut output = Vec::new();
1009        let mut writer = SizeLimitedWriter::new(&mut output, 3);
1010        writer.write_all(&[1, 2]).unwrap();
1011        let error = writer.write_all(&[3, 4]).unwrap_err().to_string();
1012        assert!(error.contains("3-byte safety limit"), "{error}");
1013        assert_eq!(output, vec![1, 2]);
1014    }
1015
1016    // ===== rebuild destructive-downgrade regression tests =====
1017
1018    use std::path::Path;
1019    use std::sync::atomic::{AtomicBool, Ordering};
1020
1021    use crate::directories::{
1022        Directory, DirectoryWriter as DirectoryWriterTrait, FileHandle, RamDirectory, RangeReadFn,
1023    };
1024    use crate::dsl::{Document, SchemaBuilder};
1025    use crate::index::{IndexConfig, IndexWriter};
1026
1027    const READ_FAIL_DOCS: usize = 5;
1028    const READ_FAIL_DIM: usize = 4;
1029    /// Flat entry layout of a single-field, flat-only `.vectors` file written
1030    /// by the segment builder (data-first format): header (16 bytes) + raw f32
1031    /// vectors + doc-id map + TOC + footer. Only the raw vector region is read
1032    /// by training collection; segment open touches the header, doc-id map,
1033    /// TOC, and footer, which all live outside this byte range.
1034    const VEC_REGION_START: u64 = 16;
1035    const VEC_REGION_END: u64 = VEC_REGION_START + (READ_FAIL_DOCS * READ_FAIL_DIM * 4) as u64;
1036
1037    /// RamDirectory wrapper whose `.vectors` handles fail range reads of the
1038    /// raw vector region while `fail_vector_reads` is armed. Segment open
1039    /// keeps succeeding, so exactly the training-collection batch reads fail —
1040    /// the I/O the rebuild path used to swallow with `if let Ok`.
1041    #[derive(Clone, Default)]
1042    struct VectorReadFailDirectory {
1043        inner: RamDirectory,
1044        fail_vector_reads: Arc<AtomicBool>,
1045        fail_all_vector_reads: Arc<AtomicBool>,
1046    }
1047
1048    #[async_trait::async_trait]
1049    impl Directory for VectorReadFailDirectory {
1050        async fn exists(&self, path: &Path) -> std::io::Result<bool> {
1051            self.inner.exists(path).await
1052        }
1053
1054        async fn file_size(&self, path: &Path) -> std::io::Result<u64> {
1055            self.inner.file_size(path).await
1056        }
1057
1058        async fn open_read(&self, path: &Path) -> std::io::Result<FileHandle> {
1059            self.inner.open_read(path).await
1060        }
1061
1062        async fn read_range(
1063            &self,
1064            path: &Path,
1065            range: std::ops::Range<u64>,
1066        ) -> std::io::Result<crate::directories::OwnedBytes> {
1067            self.inner.read_range(path, range).await
1068        }
1069
1070        async fn list_files(&self, prefix: &Path) -> std::io::Result<Vec<std::path::PathBuf>> {
1071            self.inner.list_files(prefix).await
1072        }
1073
1074        async fn open_lazy(&self, path: &Path) -> std::io::Result<FileHandle> {
1075            let handle = self.inner.open_lazy(path).await?;
1076            if path.extension().is_some_and(|ext| ext == "vectors") {
1077                let armed = Arc::clone(&self.fail_vector_reads);
1078                let fail_all = Arc::clone(&self.fail_all_vector_reads);
1079                let len = handle.len();
1080                let read_fn: RangeReadFn = Arc::new(move |range: std::ops::Range<u64>| {
1081                    let handle = handle.clone();
1082                    let armed = Arc::clone(&armed);
1083                    let fail_all = Arc::clone(&fail_all);
1084                    Box::pin(async move {
1085                        if fail_all.load(Ordering::SeqCst)
1086                            || (armed.load(Ordering::SeqCst)
1087                                && range.start >= VEC_REGION_START
1088                                && range.end <= VEC_REGION_END)
1089                        {
1090                            return Err(std::io::Error::other("injected vector data read failure"));
1091                        }
1092                        handle.read_bytes_range(range).await
1093                    })
1094                });
1095                return Ok(FileHandle::lazy(len, read_fn));
1096            }
1097            Ok(handle)
1098        }
1099    }
1100
1101    #[async_trait::async_trait]
1102    impl DirectoryWriterTrait for VectorReadFailDirectory {
1103        async fn write(&self, path: &Path, data: &[u8]) -> std::io::Result<()> {
1104            self.inner.write(path, data).await
1105        }
1106
1107        async fn delete(&self, path: &Path) -> std::io::Result<()> {
1108            self.inner.delete(path).await
1109        }
1110
1111        async fn rename(&self, from: &Path, to: &Path) -> std::io::Result<()> {
1112            self.inner.rename(from, to).await
1113        }
1114
1115        async fn sync(&self) -> std::io::Result<()> {
1116            self.inner.sync().await
1117        }
1118
1119        async fn streaming_writer(
1120            &self,
1121            path: &Path,
1122        ) -> std::io::Result<Box<dyn crate::directories::StreamingWriter>> {
1123            self.inner.streaming_writer(path).await
1124        }
1125    }
1126
1127    /// A failed read from the flat staging generation must abort training
1128    /// before artifacts or Built metadata are published.
1129    #[tokio::test]
1130    async fn build_propagates_vector_read_errors_without_publishing_artifacts() {
1131        let mut sb = SchemaBuilder::default();
1132        let embedding = sb.add_dense_vector_field_with_config(
1133            "embedding",
1134            true,
1135            true,
1136            DenseVectorConfig::ivf_tq(READ_FAIL_DIM, Some(1), 1),
1137        );
1138        let schema = sb.build();
1139
1140        let dir = VectorReadFailDirectory::default();
1141        let config = IndexConfig {
1142            merge_policy: Box::new(crate::merge::NoMergePolicy),
1143            num_indexing_threads: 1,
1144            ..Default::default()
1145        };
1146        let mut writer = IndexWriter::create(dir.clone(), schema, config)
1147            .await
1148            .unwrap();
1149        for i in 0..READ_FAIL_DOCS {
1150            let mut doc = Document::new();
1151            doc.add_dense_vector(embedding, vec![i as f32 + 1.0; READ_FAIL_DIM]);
1152            writer.add_document(doc).unwrap();
1153        }
1154        writer.commit().await.unwrap();
1155        dir.fail_vector_reads.store(true, Ordering::SeqCst);
1156        let error = writer
1157            .build_vector_index()
1158            .await
1159            .expect_err("failed sample collection must fail the build")
1160            .to_string();
1161        assert!(
1162            error.contains("injected vector data read failure"),
1163            "{error}"
1164        );
1165
1166        assert!(
1167            !writer
1168                .segment_manager
1169                .read_metadata(|meta| meta.is_field_built(embedding.0))
1170                .await,
1171            "a failed build must not publish Built metadata"
1172        );
1173        assert!(
1174            writer.segment_manager.trained().is_none(),
1175            "a failed build must not publish trained artifacts"
1176        );
1177    }
1178
1179    #[tokio::test]
1180    async fn retrain_read_failure_keeps_the_complete_published_generation() {
1181        let mut sb = SchemaBuilder::default();
1182        let embedding = sb.add_dense_vector_field_with_config(
1183            "embedding",
1184            true,
1185            true,
1186            DenseVectorConfig::ivf_tq(READ_FAIL_DIM, Some(1), 1),
1187        );
1188        let schema = sb.build();
1189        let dir = VectorReadFailDirectory::default();
1190        let config = IndexConfig {
1191            merge_policy: Box::new(crate::merge::NoMergePolicy),
1192            num_indexing_threads: 1,
1193            ..Default::default()
1194        };
1195        let mut writer = IndexWriter::create(dir.clone(), schema, config)
1196            .await
1197            .unwrap();
1198        for i in 0..READ_FAIL_DOCS {
1199            let mut doc = Document::new();
1200            doc.add_dense_vector(embedding, vec![i as f32 + 1.0; READ_FAIL_DIM]);
1201            writer.add_document(doc).unwrap();
1202        }
1203        writer.commit().await.unwrap();
1204        writer.build_vector_index().await.unwrap();
1205
1206        let old_ids = writer.segment_manager.get_segment_ids().await;
1207        let old_meta = writer
1208            .segment_manager
1209            .read_metadata(|metadata| metadata.get_field_meta(embedding.0).cloned())
1210            .await
1211            .unwrap();
1212        let old_version = writer.segment_manager.trained().unwrap().centroids[&embedding.0].version;
1213
1214        dir.fail_all_vector_reads.store(true, Ordering::SeqCst);
1215        let error = writer
1216            .retrain_vector_index()
1217            .await
1218            .expect_err("failed sample collection must abort the retrain")
1219            .to_string();
1220        assert!(
1221            error.contains("injected vector data read failure"),
1222            "{error}"
1223        );
1224        assert_eq!(writer.segment_manager.get_segment_ids().await, old_ids);
1225        assert_eq!(
1226            writer
1227                .segment_manager
1228                .read_metadata(|metadata| metadata
1229                    .get_field_meta(embedding.0)
1230                    .map(|field| (field.centroids_file.clone(), field.codebook_file.clone())))
1231                .await,
1232            Some((old_meta.centroids_file, old_meta.codebook_file)),
1233        );
1234        assert_eq!(
1235            writer.segment_manager.trained().unwrap().centroids[&embedding.0].version,
1236            old_version,
1237        );
1238    }
1239}