use std::collections::BTreeMap;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use async_trait::async_trait;
use parking_lot::RwLock;
use crate::embedding::embedder::Embedder;
use crate::error::{LaurusError, Result};
use crate::storage::Storage;
use crate::storage::prefixed::PrefixedStorage;
use crate::vector::core::distance::DistanceMetric;
use crate::vector::core::vector::Vector;
use crate::vector::index::config::VectorIndexTypeConfig;
use crate::vector::index::factory::VectorIndexFactory;
use crate::vector::index::{VectorIndex, VectorIndexStats};
use crate::vector::reader::{
SimpleVectorIterator, ValidationReport, VectorIndexMetadata, VectorIndexReader, VectorIterator,
VectorStats,
};
use crate::vector::search::searcher::{
VectorIndexQuery, VectorIndexQueryResults, VectorIndexSearcher,
};
use crate::vector::writer::VectorIndexWriter;
type FieldGroupedVectors = BTreeMap<String, Vec<(u64, String, Vector)>>;
pub(crate) const SUB_INDEX_NAME: &str = "index";
fn try_all<'a, V: 'a, F>(values: impl Iterator<Item = &'a V>, mut f: F) -> Result<()>
where
F: FnMut(&'a V) -> Result<()>,
{
let mut first_err = None;
for v in values {
if let Err(e) = f(v)
&& first_err.is_none()
{
first_err = Some(e);
}
}
first_err.map_or(Ok(()), Err)
}
fn try_all_mut<'a, V: 'a, F>(values: impl Iterator<Item = &'a mut V>, mut f: F) -> Result<()>
where
F: FnMut(&'a mut V) -> Result<()>,
{
let mut first_err = None;
for v in values {
if let Err(e) = f(v)
&& first_err.is_none()
{
first_err = Some(e);
}
}
first_err.map_or(Ok(()), Err)
}
#[derive(Debug, Clone)]
struct FieldEntry {
index: Arc<dyn VectorIndex>,
dimension: usize,
distance_metric: DistanceMetric,
}
pub struct MultiFieldVectorIndex {
fields: RwLock<BTreeMap<String, FieldEntry>>,
storage: Arc<dyn Storage>,
embedder: Arc<dyn Embedder>,
closed: AtomicBool,
}
impl std::fmt::Debug for MultiFieldVectorIndex {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MultiFieldVectorIndex")
.field("fields", &*self.fields.read())
.field("embedder", &self.embedder.name())
.field("closed", &self.is_closed())
.finish()
}
}
impl MultiFieldVectorIndex {
pub fn open_or_create(
storage: Arc<dyn Storage>,
field_configs: &BTreeMap<String, VectorIndexTypeConfig>,
embedder: Arc<dyn Embedder>,
) -> Result<Self> {
let mut fields = BTreeMap::new();
for (name, config) in field_configs {
let field_storage: Arc<dyn Storage> =
Arc::new(PrefixedStorage::new(name.clone(), storage.clone()));
let index =
VectorIndexFactory::open_or_create(field_storage, SUB_INDEX_NAME, config.clone())?;
fields.insert(
name.clone(),
FieldEntry {
dimension: config.dimension(),
distance_metric: config.distance_metric(),
index: Arc::from(index),
},
);
}
Ok(Self {
fields: RwLock::new(fields),
storage,
embedder,
closed: AtomicBool::new(false),
})
}
fn min_wal_seq(fields: &BTreeMap<String, FieldEntry>) -> u64 {
fields
.values()
.map(|f| f.index.last_wal_seq())
.min()
.unwrap_or(0)
}
}
impl VectorIndex for MultiFieldVectorIndex {
fn reader(&self) -> Result<Arc<dyn VectorIndexReader>> {
let fields = self.fields.read();
let mut readers = BTreeMap::new();
for (name, entry) in fields.iter() {
readers.insert(name.clone(), entry.index.reader()?);
}
Ok(Arc::new(MultiFieldReaderFacade::new(readers)))
}
fn writer(&self) -> Result<Box<dyn VectorIndexWriter>> {
let fields = self.fields.read();
let mut writers = BTreeMap::new();
for (name, entry) in fields.iter() {
writers.insert(name.clone(), entry.index.writer()?);
}
Ok(Box::new(MultiFieldWriter::new(writers)))
}
fn storage(&self) -> &Arc<dyn Storage> {
&self.storage
}
fn close(&self) -> Result<()> {
self.closed.store(true, Ordering::Release);
try_all(self.fields.read().values(), |f| f.index.close())
}
fn is_closed(&self) -> bool {
self.closed.load(Ordering::Acquire)
}
fn stats(&self) -> Result<VectorIndexStats> {
let fields = self.fields.read();
let mut vector_count = 0u64;
let mut total_size = 0u64;
let mut deleted_count = 0u64;
let mut last_modified = 0u64;
let mut dimension = 0usize;
for entry in fields.values() {
let s = entry.index.stats()?;
vector_count += s.vector_count;
total_size += s.total_size;
deleted_count += s.deleted_count;
last_modified = last_modified.max(s.last_modified);
if dimension == 0 {
dimension = s.dimension;
}
}
Ok(VectorIndexStats {
vector_count,
dimension,
total_size,
deleted_count,
last_modified,
})
}
fn optimize(&self) -> Result<()> {
try_all(self.fields.read().values(), |f| f.index.optimize())
}
fn refresh(&self) -> Result<()> {
try_all(self.fields.read().values(), |f| f.index.refresh())
}
fn retain_writer_after_commit(&self) -> bool {
let fields = self.fields.read();
!fields.is_empty()
&& fields
.values()
.all(|f| f.index.retain_writer_after_commit())
}
fn searcher(&self) -> Result<Box<dyn VectorIndexSearcher>> {
let fields = self.fields.read();
let mut searchers = BTreeMap::new();
for (name, entry) in fields.iter() {
searchers.insert(
name.clone(),
SearcherEntry {
searcher: entry.index.searcher()?,
dimension: entry.dimension,
distance_metric: entry.distance_metric,
},
);
}
Ok(Box::new(MultiFieldFanoutSearcher::new(searchers)))
}
fn embedder(&self) -> Arc<dyn Embedder> {
self.embedder.clone()
}
fn last_wal_seq(&self) -> u64 {
Self::min_wal_seq(&self.fields.read())
}
fn set_last_wal_seq(&self, seq: u64) -> Result<()> {
try_all(self.fields.read().values(), |f| {
f.index.set_last_wal_seq(seq)
})
}
fn supports_soft_delete(&self) -> bool {
let fields = self.fields.read();
!fields.is_empty() && fields.values().all(|f| f.index.supports_soft_delete())
}
fn soft_delete_document(&self, doc_id: u64) -> Result<()> {
try_all(self.fields.read().values(), |f| {
f.index.soft_delete_document(doc_id)
})
}
fn persist_deletions(&self) -> Result<()> {
try_all(self.fields.read().values(), |f| f.index.persist_deletions())
}
fn maybe_auto_compact(&self) -> Result<bool> {
let fields = self.fields.read();
let mut compacted_any = false;
let mut first_err = None;
for f in fields.values() {
match f.index.maybe_auto_compact() {
Ok(true) => compacted_any = true,
Ok(false) => {}
Err(e) if first_err.is_none() => first_err = Some(e),
Err(_) => {}
}
}
match first_err {
Some(e) => Err(e),
None => Ok(compacted_any),
}
}
fn supports_dynamic_fields(&self) -> bool {
true
}
fn add_field(&self, name: &str, config: VectorIndexTypeConfig) -> Result<()> {
let mut fields = self.fields.write();
if fields.contains_key(name) {
return Err(LaurusError::invalid_argument(format!(
"vector field '{name}' already exists"
)));
}
let field_storage: Arc<dyn Storage> =
Arc::new(PrefixedStorage::new(name.to_string(), self.storage.clone()));
let index =
VectorIndexFactory::open_or_create(field_storage, SUB_INDEX_NAME, config.clone())?;
let current_min = Self::min_wal_seq(&fields);
index.set_last_wal_seq(current_min)?;
fields.insert(
name.to_string(),
FieldEntry {
dimension: config.dimension(),
distance_metric: config.distance_metric(),
index: Arc::from(index),
},
);
Ok(())
}
fn remove_field(&self, name: &str) -> Result<()> {
self.fields.write().remove(name);
Ok(())
}
fn field_dimensions(&self) -> BTreeMap<String, usize> {
self.fields
.read()
.iter()
.map(|(name, entry)| (name.clone(), entry.dimension))
.collect()
}
}
#[derive(Debug)]
struct MultiFieldWriter {
writers: BTreeMap<String, Box<dyn VectorIndexWriter>>,
closed: bool,
}
impl MultiFieldWriter {
fn new(writers: BTreeMap<String, Box<dyn VectorIndexWriter>>) -> Self {
Self {
writers,
closed: false,
}
}
fn group_by_field(&self, vectors: Vec<(u64, String, Vector)>) -> Result<FieldGroupedVectors> {
let mut grouped: FieldGroupedVectors = BTreeMap::new();
for (doc_id, field_name, vector) in vectors {
if !self.writers.contains_key(&field_name) {
return Err(LaurusError::invalid_argument(format!(
"unknown vector field '{field_name}': no index configured for it"
)));
}
grouped
.entry(field_name.clone())
.or_default()
.push((doc_id, field_name, vector));
}
Ok(grouped)
}
}
#[async_trait]
impl VectorIndexWriter for MultiFieldWriter {
fn next_vector_id(&self) -> u64 {
self.writers
.values()
.map(|w| w.next_vector_id())
.max()
.unwrap_or(0)
}
fn build(&mut self, vectors: Vec<(u64, String, Vector)>) -> Result<()> {
self.add_vectors(vectors)
}
fn add_vectors(&mut self, vectors: Vec<(u64, String, Vector)>) -> Result<()> {
let grouped = self.group_by_field(vectors)?;
for (field_name, field_vectors) in grouped {
let writer = self
.writers
.get_mut(&field_name)
.expect("field name validated by group_by_field");
writer.add_vectors(field_vectors)?;
}
Ok(())
}
fn finalize(&mut self) -> Result<()> {
try_all_mut(self.writers.values_mut(), |w| w.finalize())
}
fn progress(&self) -> f32 {
if self.writers.is_empty() {
return 1.0;
}
self.writers.values().map(|w| w.progress()).sum::<f32>() / self.writers.len() as f32
}
fn estimated_memory_usage(&self) -> usize {
self.writers
.values()
.map(|w| w.estimated_memory_usage())
.sum()
}
fn vectors(&self) -> &[(u64, String, Vector)] {
&[]
}
fn write(&self) -> Result<()> {
try_all(self.writers.values(), |w| w.write())
}
fn has_storage(&self) -> bool {
self.writers.values().all(|w| w.has_storage())
}
fn delete_document(&mut self, doc_id: u64) -> Result<()> {
try_all_mut(self.writers.values_mut(), |w| w.delete_document(doc_id))
}
fn has_pending_changes(&self) -> bool {
self.writers.values().any(|w| w.has_pending_changes())
}
fn delete_documents(&mut self, field: &str, value: &str) -> Result<usize> {
let mut total = 0usize;
let mut first_err = None;
for w in self.writers.values_mut() {
match w.delete_documents(field, value) {
Ok(n) => total += n,
Err(e) if first_err.is_none() => first_err = Some(e),
Err(_) => {}
}
}
match first_err {
Some(e) => Err(e),
None => Ok(total),
}
}
fn commit(&mut self) -> Result<()> {
let mut first_err = None;
for w in self.writers.values_mut() {
if let Err(e) = w.commit()
&& first_err.is_none()
{
first_err = Some(e);
}
}
match first_err {
Some(e) => Err(e),
None => Ok(()),
}
}
async fn add_value(
&mut self,
doc_id: u64,
field_name: String,
value: crate::data::DataValue,
) -> Result<()> {
let writer = self.writers.get_mut(&field_name).ok_or_else(|| {
LaurusError::invalid_argument(format!(
"unknown vector field '{field_name}': no index configured for it"
))
})?;
writer.add_value(doc_id, field_name.clone(), value).await
}
fn rollback(&mut self) -> Result<()> {
try_all_mut(self.writers.values_mut(), |w| w.rollback())
}
fn pending_docs(&self) -> u64 {
self.writers
.values()
.map(|w| w.pending_docs())
.max()
.unwrap_or(0)
}
fn close(&mut self) -> Result<()> {
if self.closed {
return Ok(());
}
let result = try_all_mut(self.writers.values_mut(), |w| w.close());
self.closed = true;
result
}
fn is_closed(&self) -> bool {
self.closed || self.writers.values().all(|w| w.is_closed())
}
fn optimize(&mut self) -> Result<()> {
try_all_mut(self.writers.values_mut(), |w| w.optimize())
}
fn build_reader(&self) -> Result<Arc<dyn VectorIndexReader>> {
let mut readers = BTreeMap::new();
for (name, w) in &self.writers {
readers.insert(name.clone(), w.build_reader()?);
}
Ok(Arc::new(MultiFieldReaderFacade::new(readers)))
}
}
#[derive(Debug)]
struct SearcherEntry {
searcher: Box<dyn VectorIndexSearcher>,
dimension: usize,
distance_metric: DistanceMetric,
}
#[derive(Debug)]
struct MultiFieldFanoutSearcher {
fields: BTreeMap<String, SearcherEntry>,
}
impl MultiFieldFanoutSearcher {
fn new(fields: BTreeMap<String, SearcherEntry>) -> Self {
Self { fields }
}
fn resolve_field(&self, field_name: &str) -> Result<&SearcherEntry> {
self.fields.get(field_name).ok_or_else(|| {
LaurusError::invalid_argument(format!(
"unknown vector field '{field_name}': no index configured for it"
))
})
}
fn candidates_for_dimension(&self, dim: usize) -> Vec<&SearcherEntry> {
self.fields
.values()
.filter(|e| e.dimension == dim)
.collect()
}
fn search_fanout(&self, request: &VectorIndexQuery) -> Result<VectorIndexQueryResults> {
let dim = request.query.data.len();
let candidates = self.candidates_for_dimension(dim);
if candidates.is_empty() {
return Ok(VectorIndexQueryResults::new());
}
if candidates.len() == 1 {
return candidates[0].searcher.search(request);
}
let homogeneous_metric = candidates
.windows(2)
.all(|w| w[0].distance_metric == w[1].distance_metric);
let mut merged = VectorIndexQueryResults::new();
let mut candidates_examined = 0usize;
let mut search_time_ms = 0f64;
for entry in &candidates {
let r = entry.searcher.search(request)?;
candidates_examined += r.candidates_examined;
search_time_ms += r.search_time_ms;
merged.results.extend(r.results);
merged.query_metadata.extend(r.query_metadata);
}
merged.candidates_examined = candidates_examined;
merged.search_time_ms = search_time_ms;
if homogeneous_metric {
merged.sort_by_distance();
} else {
merged.sort_by_similarity();
}
merged.take_top_k(request.params.top_k);
Ok(merged)
}
}
impl VectorIndexSearcher for MultiFieldFanoutSearcher {
fn search(&self, request: &VectorIndexQuery) -> Result<VectorIndexQueryResults> {
match &request.field_name {
Some(field_name) => self.resolve_field(field_name)?.searcher.search(request),
None => self.search_fanout(request),
}
}
fn count(&self, request: VectorIndexQuery) -> Result<u64> {
if let Some(field_name) = request.field_name.clone() {
return self.resolve_field(&field_name)?.searcher.count(request);
}
let dim = request.query.data.len();
let mut total = 0u64;
for entry in self.candidates_for_dimension(dim) {
total += entry.searcher.count(request.clone())?;
}
Ok(total)
}
fn warmup(&mut self) -> Result<()> {
for entry in self.fields.values_mut() {
entry.searcher.warmup()?;
}
Ok(())
}
fn parallel_threshold(&self) -> usize {
self.fields
.values()
.map(|e| e.searcher.parallel_threshold())
.min()
.unwrap_or(4)
}
fn search_batch_with_threshold(
&self,
queries: &[VectorIndexQuery],
parallel_threshold: usize,
) -> Result<Vec<VectorIndexQueryResults>> {
let mut buckets: BTreeMap<Option<String>, Vec<usize>> = BTreeMap::new();
for (i, q) in queries.iter().enumerate() {
buckets.entry(q.field_name.clone()).or_default().push(i);
}
let mut results: Vec<Option<VectorIndexQueryResults>> =
(0..queries.len()).map(|_| None).collect();
for (field_name, indices) in buckets {
match field_name {
Some(field_name) => {
let entry = self.resolve_field(&field_name)?;
let bucket_queries: Vec<VectorIndexQuery> =
indices.iter().map(|&i| queries[i].clone()).collect();
let bucket_results = entry
.searcher
.search_batch_with_threshold(&bucket_queries, parallel_threshold)?;
for (i, r) in indices.into_iter().zip(bucket_results) {
results[i] = Some(r);
}
}
None => {
for i in indices {
results[i] = Some(self.search_fanout(&queries[i])?);
}
}
}
}
Ok(results
.into_iter()
.map(|r| r.expect("every query index is populated by exactly one bucket above"))
.collect())
}
}
#[derive(Debug)]
struct MultiFieldReaderFacade {
readers: BTreeMap<String, Arc<dyn VectorIndexReader>>,
}
impl MultiFieldReaderFacade {
fn new(readers: BTreeMap<String, Arc<dyn VectorIndexReader>>) -> Self {
Self { readers }
}
}
impl VectorIndexReader for MultiFieldReaderFacade {
fn as_any(&self) -> &dyn std::any::Any {
self
}
fn get_vector(&self, doc_id: u64, field_name: &str) -> Result<Option<Vector>> {
match self.readers.get(field_name) {
Some(r) => r.get_vector(doc_id, field_name),
None => Ok(None),
}
}
fn get_vectors_for_doc(&self, doc_id: u64) -> Result<Vec<(String, Vector)>> {
let mut out = Vec::new();
for r in self.readers.values() {
out.extend(r.get_vectors_for_doc(doc_id)?);
}
Ok(out)
}
fn get_vectors(&self, doc_ids: &[(u64, String)]) -> Result<Vec<Option<Vector>>> {
doc_ids
.iter()
.map(|(doc_id, field_name)| self.get_vector(*doc_id, field_name))
.collect()
}
fn vector_ids(&self) -> Result<Vec<(u64, String)>> {
let mut out = Vec::new();
for r in self.readers.values() {
out.extend(r.vector_ids()?);
}
Ok(out)
}
fn doc_ids_for_field(&self, field_name: &str) -> Arc<[u64]> {
match self.readers.get(field_name) {
Some(r) => r.doc_ids_for_field(field_name),
None => Vec::new().into(),
}
}
fn vector_count(&self) -> usize {
self.readers.values().map(|r| r.vector_count()).sum()
}
fn dimension(&self) -> usize {
self.readers
.values()
.next()
.map(|r| r.dimension())
.unwrap_or(0)
}
fn distance_metric(&self) -> DistanceMetric {
self.readers
.values()
.next()
.map(|r| r.distance_metric())
.unwrap_or(DistanceMetric::Cosine)
}
fn stats(&self) -> VectorStats {
let mut vector_count = 0;
let mut memory_usage = 0;
let mut build_time_ms = 0;
for r in self.readers.values() {
let s = r.stats();
vector_count += s.vector_count;
memory_usage += s.memory_usage;
build_time_ms = build_time_ms.max(s.build_time_ms);
}
VectorStats {
vector_count,
dimension: self.dimension(),
memory_usage,
build_time_ms,
}
}
fn contains_vector(&self, doc_id: u64, field_name: &str) -> bool {
self.readers
.get(field_name)
.map(|r| r.contains_vector(doc_id, field_name))
.unwrap_or(false)
}
fn get_vector_range(
&self,
start_doc_id: u64,
end_doc_id: u64,
) -> Result<Vec<(u64, String, Vector)>> {
let mut out = Vec::new();
for r in self.readers.values() {
out.extend(r.get_vector_range(start_doc_id, end_doc_id)?);
}
Ok(out)
}
fn get_vectors_by_field(&self, field_name: &str) -> Result<Vec<(u64, Vector)>> {
match self.readers.get(field_name) {
Some(r) => r.get_vectors_by_field(field_name),
None => Ok(Vec::new()),
}
}
fn field_names(&self) -> Result<Vec<String>> {
Ok(self.readers.keys().cloned().collect())
}
fn vector_iterator(&self) -> Result<Box<dyn VectorIterator>> {
let mut all = Vec::new();
for r in self.readers.values() {
let mut it = r.vector_iterator()?;
while let Some(item) = it.next()? {
all.push(item);
}
}
Ok(Box::new(SimpleVectorIterator::new(all)))
}
fn metadata(&self) -> Result<VectorIndexMetadata> {
Ok(VectorIndexMetadata {
index_type: "MultiField".to_string(),
created_at: chrono::Utc::now(),
modified_at: chrono::Utc::now(),
version: "1.0".to_string(),
build_config: serde_json::json!({
"fields": self.readers.keys().cloned().collect::<Vec<_>>(),
}),
custom_metadata: std::collections::HashMap::new(),
})
}
fn validate(&self) -> Result<ValidationReport> {
let mut errors = Vec::new();
let mut warnings = Vec::new();
let mut repair_suggestions = Vec::new();
for (name, r) in &self.readers {
let report = r.validate()?;
errors.extend(report.errors.into_iter().map(|e| format!("[{name}] {e}")));
warnings.extend(report.warnings.into_iter().map(|w| format!("[{name}] {w}")));
repair_suggestions.extend(report.repair_suggestions);
}
Ok(ValidationReport {
is_valid: errors.is_empty(),
errors,
warnings,
repair_suggestions,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicUsize, Ordering};
#[derive(Debug)]
struct FailingWriter {
commit_attempts: Arc<AtomicUsize>,
}
#[async_trait]
impl VectorIndexWriter for FailingWriter {
fn next_vector_id(&self) -> u64 {
0
}
fn build(&mut self, vectors: Vec<(u64, String, Vector)>) -> Result<()> {
self.add_vectors(vectors)
}
fn add_vectors(&mut self, _vectors: Vec<(u64, String, Vector)>) -> Result<()> {
Ok(())
}
fn finalize(&mut self) -> Result<()> {
Ok(())
}
fn progress(&self) -> f32 {
1.0
}
fn estimated_memory_usage(&self) -> usize {
0
}
fn vectors(&self) -> &[(u64, String, Vector)] {
&[]
}
fn write(&self) -> Result<()> {
Ok(())
}
fn has_storage(&self) -> bool {
true
}
fn delete_document(&mut self, _doc_id: u64) -> Result<()> {
Ok(())
}
fn commit(&mut self) -> Result<()> {
self.commit_attempts.fetch_add(1, Ordering::SeqCst);
Err(LaurusError::internal("simulated commit failure"))
}
fn rollback(&mut self) -> Result<()> {
Ok(())
}
fn pending_docs(&self) -> u64 {
0
}
fn close(&mut self) -> Result<()> {
Ok(())
}
fn is_closed(&self) -> bool {
false
}
fn build_reader(&self) -> Result<Arc<dyn VectorIndexReader>> {
Ok(Arc::new(MultiFieldReaderFacade::new(BTreeMap::new())))
}
}
#[derive(Debug)]
struct SucceedingWriter {
commit_attempts: Arc<AtomicUsize>,
}
#[async_trait]
impl VectorIndexWriter for SucceedingWriter {
fn next_vector_id(&self) -> u64 {
0
}
fn build(&mut self, vectors: Vec<(u64, String, Vector)>) -> Result<()> {
self.add_vectors(vectors)
}
fn add_vectors(&mut self, _vectors: Vec<(u64, String, Vector)>) -> Result<()> {
Ok(())
}
fn finalize(&mut self) -> Result<()> {
Ok(())
}
fn progress(&self) -> f32 {
1.0
}
fn estimated_memory_usage(&self) -> usize {
0
}
fn vectors(&self) -> &[(u64, String, Vector)] {
&[]
}
fn write(&self) -> Result<()> {
Ok(())
}
fn has_storage(&self) -> bool {
true
}
fn delete_document(&mut self, _doc_id: u64) -> Result<()> {
Ok(())
}
fn commit(&mut self) -> Result<()> {
self.commit_attempts.fetch_add(1, Ordering::SeqCst);
Ok(())
}
fn rollback(&mut self) -> Result<()> {
Ok(())
}
fn pending_docs(&self) -> u64 {
0
}
fn close(&mut self) -> Result<()> {
Ok(())
}
fn is_closed(&self) -> bool {
false
}
fn build_reader(&self) -> Result<Arc<dyn VectorIndexReader>> {
Ok(Arc::new(MultiFieldReaderFacade::new(BTreeMap::new())))
}
}
#[test]
fn commit_always_attempts_every_field_and_never_poisons() {
let good_commits = Arc::new(AtomicUsize::new(0));
let bad_commits = Arc::new(AtomicUsize::new(0));
let mut writers: BTreeMap<String, Box<dyn VectorIndexWriter>> = BTreeMap::new();
writers.insert(
"good".to_string(),
Box::new(SucceedingWriter {
commit_attempts: good_commits.clone(),
}),
);
writers.insert(
"bad".to_string(),
Box::new(FailingWriter {
commit_attempts: bad_commits.clone(),
}),
);
let mut writer = MultiFieldWriter::new(writers);
let err = writer.commit().unwrap_err();
assert!(format!("{err:?}").contains("simulated commit failure"));
assert_eq!(good_commits.load(Ordering::SeqCst), 1);
assert_eq!(bad_commits.load(Ordering::SeqCst), 1);
writer
.add_vectors(vec![(1, "good".to_string(), Vector::new(vec![0.0]))])
.unwrap();
writer.commit().unwrap_err();
assert_eq!(good_commits.load(Ordering::SeqCst), 2);
assert_eq!(bad_commits.load(Ordering::SeqCst), 2);
}
}