1use 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
22const MAX_IVF_CLUSTERS: usize = 1_048_576;
25const MIN_TRAINING_POINTS_PER_CENTROID: usize = 39;
28const COARSE_TRAINING_POINTS_PER_CENTROID: usize = 256;
31const MAX_SAMPLE_READ_BYTES: usize = 64 * 1024 * 1024;
34const SAMPLE_BLOCK: usize = 256;
35const 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 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
116struct 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#[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 pub async fn build_vector_index(&self) -> Result<()> {
265 self.build_vector_generation(VectorGenerationMode::BuildMissing)
266 .await
267 }
268
269 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 drop(snapshot);
380 drop(artifact_update);
381
382 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 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 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 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 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 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 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 .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 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 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 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 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 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 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(¢roids, &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 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 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 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 #[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 #[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}