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