use std::sync::Arc;
#[cfg(not(target_arch = "wasm32"))]
use rayon::prelude::*;
use crate::error::{LaurusError, Result};
use crate::storage::Storage;
use crate::vector::core::quantization::ScalarQuantParams;
use crate::vector::core::vector::Vector;
use crate::vector::index::FlatIndexConfig;
use crate::vector::index::alloc_bounds::checked_capacity;
use crate::vector::index::field::LegacyVectorFieldWriter;
use crate::vector::index::format::{
QuantHeader, VERSION_FIELD_DICT, VectorSegmentHeader, build_field_dict, record_prefix_size,
};
use crate::vector::index::quantized_io::{
quantize_segment, quantized_record_payload_size, read_dequantized_vector,
write_quantized_record,
};
use crate::vector::writer::{VectorIndexWriter, VectorIndexWriterConfig};
#[derive(Debug)]
pub struct FlatIndexWriter {
index_config: FlatIndexConfig,
writer_config: VectorIndexWriterConfig,
storage: Option<Arc<dyn Storage>>,
path: String,
vectors: Vec<(u64, String, Vector)>,
is_finalized: bool,
total_vectors_to_add: Option<usize>,
next_vec_id: u64,
}
impl FlatIndexWriter {
pub fn new(
index_config: FlatIndexConfig,
writer_config: VectorIndexWriterConfig,
path: impl Into<String>,
) -> Result<Self> {
Ok(Self {
index_config,
writer_config,
storage: None,
path: path.into(),
vectors: Vec::new(),
is_finalized: false,
total_vectors_to_add: None,
next_vec_id: 0,
})
}
pub fn with_storage(
index_config: FlatIndexConfig,
writer_config: VectorIndexWriterConfig,
path: impl Into<String>,
storage: Arc<dyn Storage>,
) -> Result<Self> {
let path = path.into();
let file_name = format!("{}.flat", path);
if storage.file_exists(&file_name) {
return Self::load(index_config, writer_config, storage, &path);
}
Ok(Self {
index_config,
writer_config,
storage: Some(storage),
path,
vectors: Vec::new(),
is_finalized: false,
total_vectors_to_add: None,
next_vec_id: 0,
})
}
pub fn into_field_writer(self, field_name: impl Into<String>) -> LegacyVectorFieldWriter<Self> {
LegacyVectorFieldWriter::new(field_name, self)
}
pub fn load(
index_config: FlatIndexConfig,
writer_config: VectorIndexWriterConfig,
storage: Arc<dyn Storage>,
path: &str,
) -> Result<Self> {
use std::io::{Read, Seek};
let file_name = format!("{}.flat", path);
let mut input = storage.open_input(&file_name)?;
let file_size = input.size()?;
let mut num_vectors_buf = [0u8; 4];
input.read_exact(&mut num_vectors_buf)?;
let num_vectors = u32::from_le_bytes(num_vectors_buf) as usize;
let mut dimension_buf = [0u8; 4];
input.read_exact(&mut dimension_buf)?;
let dimension = u32::from_le_bytes(dimension_buf) as usize;
if dimension != index_config.dimension {
return Err(LaurusError::InvalidOperation(format!(
"Dimension mismatch: expected {}, found {}",
index_config.dimension, dimension
)));
}
let header_available =
file_size.saturating_sub(input.stream_position().map_err(LaurusError::Io)?);
let header = VectorSegmentHeader::read_from(&mut input, header_available)?;
let params = match &header.quant {
QuantHeader::Scalar8Bit(p) => *p,
QuantHeader::ProductQuantization { .. } => {
return Err(crate::error::LaurusError::NotImplemented(
"Product quantization (Issue #481 Stage 3) is HNSW-only; \
the Flat writer does not support PQ segments yet"
.to_string(),
));
}
#[cfg(feature = "pq-fastscan")]
QuantHeader::ProductQuantizationFastScan { .. } => {
return Err(crate::error::LaurusError::NotImplemented(
"PQ FastScan (#695) is HNSW-only; the Flat writer does not \
support PQ FastScan segments"
.to_string(),
));
}
};
let records_remaining =
file_size.saturating_sub(input.stream_position().map_err(LaurusError::Io)?);
let record_stride =
record_prefix_size(header.version) + quantized_record_payload_size(dimension) as u64;
checked_capacity(
num_vectors,
record_stride,
records_remaining,
"flat num_vectors",
)?;
let mut vectors = Vec::with_capacity(num_vectors);
for _ in 0..num_vectors {
let mut doc_id_buf = [0u8; 8];
input.read_exact(&mut doc_id_buf)?;
let doc_id = u64::from_le_bytes(doc_id_buf);
let field_name =
header.read_record_field(&mut input, records_remaining, "flat field_name_len")?;
let values = read_dequantized_vector(&mut input, dimension, ¶ms)?;
vectors.push((doc_id, field_name, Vector::new(values)));
}
let max_id = vectors.iter().map(|(id, _, _)| *id).max().unwrap_or(0);
let next_vec_id = if num_vectors > 0 { max_id + 1 } else { 0 };
Ok(Self {
index_config,
writer_config,
storage: Some(storage),
path: path.to_string(),
vectors,
is_finalized: true,
total_vectors_to_add: Some(num_vectors),
next_vec_id,
})
}
pub fn set_expected_vector_count(&mut self, count: usize) {
self.total_vectors_to_add = Some(count);
}
pub fn vectors(&self) -> &[(u64, String, Vector)] {
&self.vectors
}
fn validate_vectors(&self, vectors: &[(u64, String, Vector)]) -> Result<()> {
if vectors.is_empty() {
return Ok(());
}
for (doc_id, _field_name, vector) in vectors {
if vector.dimension() != self.index_config.dimension {
return Err(LaurusError::InvalidOperation(format!(
"Vector {} has dimension {}, expected {}",
doc_id,
vector.dimension(),
self.index_config.dimension
)));
}
if !vector.is_valid() {
return Err(LaurusError::InvalidOperation(format!(
"Vector {doc_id} contains invalid values (NaN or infinity)"
)));
}
}
Ok(())
}
fn normalize_vectors(&self, vectors: &mut [(u64, String, Vector)]) {
if !self.index_config.normalize_vectors {
return;
}
#[cfg(not(target_arch = "wasm32"))]
if self.writer_config.parallel_build && vectors.len() > 100 {
vectors.par_iter_mut().for_each(|(_, _, vector)| {
vector.normalize();
});
return;
}
for (_, _, vector) in vectors {
vector.normalize();
}
}
fn check_memory_limit(&self) -> Result<()> {
if let Some(limit) = self.writer_config.memory_limit {
let current_usage = self.estimated_memory_usage();
if current_usage > limit {
return Err(LaurusError::ResourceExhausted(format!(
"Memory usage {current_usage} bytes exceeds limit {limit} bytes"
)));
}
}
Ok(())
}
fn sort_vectors(&mut self) {
#[cfg(not(target_arch = "wasm32"))]
if self.writer_config.parallel_build && self.vectors.len() as u64 > 10000 {
self.vectors
.par_sort_by(|(doc_id_a, field_a, _), (doc_id_b, field_b, _)| {
doc_id_a.cmp(doc_id_b).then_with(|| field_a.cmp(field_b))
});
return;
}
self.vectors
.sort_by(|(doc_id_a, field_a, _), (doc_id_b, field_b, _)| {
doc_id_a.cmp(doc_id_b).then_with(|| field_a.cmp(field_b))
});
}
fn deduplicate_vectors(&mut self) {
if self.vectors.is_empty() {
return;
}
self.sort_vectors();
let mut unique_vectors = Vec::new();
let mut last_key: Option<(u64, String)> = None;
for (doc_id, field_name, vector) in std::mem::take(&mut self.vectors) {
let current_key = (doc_id, field_name.clone());
if last_key.as_ref() != Some(¤t_key) {
unique_vectors.push((doc_id, field_name, vector));
last_key = Some(current_key);
} else {
if let Some((_, _, last_vector)) = unique_vectors.last_mut() {
*last_vector = vector;
}
}
}
self.vectors = unique_vectors;
}
}
#[async_trait::async_trait]
impl VectorIndexWriter for FlatIndexWriter {
fn next_vector_id(&self) -> u64 {
self.next_vec_id
}
fn build(&mut self, mut vectors: Vec<(u64, String, Vector)>) -> Result<()> {
if self.is_finalized {
self.is_finalized = false;
}
self.validate_vectors(&vectors)?;
self.normalize_vectors(&mut vectors);
if let Some(max_id) = vectors.iter().map(|(id, _, _)| *id).max()
&& max_id >= self.next_vec_id
{
self.next_vec_id = max_id + 1;
}
self.vectors = vectors;
self.total_vectors_to_add = Some(self.vectors.len());
self.check_memory_limit()?;
Ok(())
}
fn add_vectors(&mut self, mut vectors: Vec<(u64, String, Vector)>) -> Result<()> {
if self.is_finalized {
self.is_finalized = false;
}
self.validate_vectors(&vectors)?;
self.normalize_vectors(&mut vectors);
if let Some(max_id) = vectors.iter().map(|(id, _, _)| *id).max()
&& max_id >= self.next_vec_id
{
self.next_vec_id = max_id + 1;
}
self.vectors.extend(vectors);
self.check_memory_limit()?;
Ok(())
}
fn finalize(&mut self) -> Result<()> {
if self.is_finalized {
return Ok(());
}
self.deduplicate_vectors();
self.sort_vectors();
self.is_finalized = true;
Ok(())
}
fn progress(&self) -> f32 {
if let Some(total) = self.total_vectors_to_add {
if total == 0 {
if self.is_finalized { 1.0 } else { 0.0 }
} else {
let current = self.vectors.len() as u64 as f32;
let progress = current / total as f32;
if self.is_finalized {
1.0
} else {
progress.min(0.99) }
}
} else if self.is_finalized {
1.0
} else {
0.0
}
}
fn estimated_memory_usage(&self) -> usize {
let vector_memory = self.vectors.len()
* (
8 + self.index_config.dimension * 4 + std::mem::size_of::<Vector>()
);
let metadata_memory = self.vectors.len() * 64;
vector_memory + metadata_memory
}
fn vectors(&self) -> &[(u64, String, Vector)] {
&self.vectors
}
fn write(&self) -> Result<()> {
use std::io::Write;
if !self.is_finalized {
return Err(LaurusError::InvalidOperation(
"Index must be finalized before writing".to_string(),
));
}
let storage = self
.storage
.as_ref()
.ok_or_else(|| LaurusError::InvalidOperation("No storage configured".to_string()))?;
let file_name = format!("{}.flat", self.path);
let tmp_name = format!("{}.flat.tmp", self.path);
let mut output = storage.create_output(&tmp_name)?;
let vector_count: u32 = self.vectors.len().try_into().map_err(|_| {
LaurusError::InvalidOperation(format!(
"Vector count {} exceeds u32::MAX",
self.vectors.len()
))
})?;
output.write_all(&vector_count.to_le_bytes())?;
output.write_all(&(self.index_config.dimension as u32).to_le_bytes())?;
let f32_vectors: Vec<Vector> = self.vectors.iter().map(|(_, _, v)| v.clone()).collect();
let (params, records) = if f32_vectors.is_empty() {
(
ScalarQuantParams {
offset: 0.0,
scale: 1.0,
},
Vec::new(),
)
} else {
quantize_segment(&f32_vectors, self.index_config.dimension)?
};
let (field_dict, field_ids) =
build_field_dict(self.vectors.iter().map(|(_, f, _)| f.as_str()))?;
VectorSegmentHeader::scalar_8bit(params)
.with_version(VERSION_FIELD_DICT)
.with_field_dict(field_dict)
.write_to(&mut output)?;
for ((doc_id, field_name, _), (int8, meta)) in self.vectors.iter().zip(records.iter()) {
output.write_all(&doc_id.to_le_bytes())?;
output.write_all(&field_ids[field_name.as_str()].to_le_bytes())?;
write_quantized_record(&mut output, int8, *meta)?;
}
output.close()?;
storage.rename_file(&tmp_name, &file_name)?;
if let Some(rerank_kind) = self.index_config.rerank_storage {
let sidecar_name = format!("{}.f32", file_name);
let sidecar_tmp = format!("{}.f32.tmp", file_name);
let mut sidecar_out = storage.create_output(&sidecar_tmp)?;
let mut payload: Vec<f32> =
Vec::with_capacity(self.vectors.len() * self.index_config.dimension);
for (_, _, v) in &self.vectors {
payload.extend_from_slice(&v.data);
}
crate::vector::index::rerank_sidecar::write_sidecar(
&mut sidecar_out,
rerank_kind,
self.index_config.dimension as u32,
&payload,
)?;
sidecar_out.flush()?;
drop(sidecar_out);
storage.rename_file(&sidecar_tmp, &sidecar_name)?;
}
Ok(())
}
fn has_storage(&self) -> bool {
self.storage.is_some()
}
fn delete_document(&mut self, doc_id: u64) -> Result<()> {
if self.is_finalized {
self.is_finalized = false;
}
self.vectors.retain(|(id, _, _)| *id != doc_id);
Ok(())
}
fn delete_documents(&mut self, _field: &str, _value: &str) -> Result<usize> {
if self.is_finalized {
return Err(LaurusError::InvalidOperation(
"Cannot delete documents from finalized index".to_string(),
));
}
Ok(0)
}
fn rollback(&mut self) -> Result<()> {
self.vectors.clear();
self.is_finalized = false;
self.next_vec_id = 0;
Ok(())
}
fn pending_docs(&self) -> u64 {
if self.is_finalized {
0
} else {
self.vectors.len() as u64
}
}
fn close(&mut self) -> Result<()> {
self.vectors.clear();
self.is_finalized = true;
Ok(())
}
fn is_closed(&self) -> bool {
self.is_finalized && self.vectors.is_empty()
}
fn build_reader(&self) -> Result<Arc<dyn crate::vector::reader::VectorIndexReader>> {
use crate::vector::index::flat::reader::FlatVectorIndexReader;
let storage = self.storage.as_ref().ok_or_else(|| {
LaurusError::InvalidOperation("Cannot build reader: storage not configured".to_string())
})?;
let reader = FlatVectorIndexReader::load(
storage.clone(),
&self.path,
self.index_config.distance_metric,
)?;
Ok(Arc::new(reader))
}
}