use std::any::Any;
use std::collections::BTreeMap;
use std::sync::Arc;
use async_trait::async_trait;
use laurus::storage::Storage;
use laurus::storage::memory::{MemoryStorage, MemoryStorageConfig};
use laurus::vector::index::VectorIndex;
use laurus::vector::index::config::{HnswIndexConfig, VectorIndexTypeConfig};
use laurus::vector::index::multi_field::MultiFieldVectorIndex;
use laurus::vector::search::searcher::{VectorIndexQuery, VectorIndexQueryParams};
use laurus::vector::{DistanceMetric, Vector};
use laurus::{EmbedInput, EmbedInputType, Embedder, LaurusError, Result};
#[derive(Debug)]
struct MockEmbedder;
#[async_trait]
impl Embedder for MockEmbedder {
async fn embed(&self, _input: &EmbedInput<'_>) -> Result<Vector> {
Err(LaurusError::invalid_argument(
"embedding not used by this test",
))
}
fn supported_input_types(&self) -> Vec<EmbedInputType> {
vec![EmbedInputType::Text]
}
fn name(&self) -> &str {
"mock"
}
fn as_any(&self) -> &dyn Any {
self
}
}
fn storage() -> Arc<dyn Storage> {
Arc::new(MemoryStorage::new(MemoryStorageConfig::default()))
}
fn hnsw_config(dimension: usize, distance_metric: DistanceMetric) -> VectorIndexTypeConfig {
VectorIndexTypeConfig::HNSW(HnswIndexConfig {
dimension,
distance_metric,
m: 16,
ef_construction: 100,
normalize_vectors: distance_metric == DistanceMetric::Cosine,
..Default::default()
})
}
fn vec_of(values: &[f32]) -> Vector {
Vector::new(values.to_vec())
}
fn query(v: Vector, field_name: Option<&str>, top_k: usize) -> VectorIndexQuery {
VectorIndexQuery {
query: v,
params: VectorIndexQueryParams {
top_k,
..Default::default()
},
field_name: field_name.map(str::to_string),
filter: None,
}
}
#[test]
fn two_fields_same_doc_id_both_retained() {
let mut fields = BTreeMap::new();
fields.insert(
"title_vec".to_string(),
hnsw_config(4, DistanceMetric::Cosine),
);
fields.insert(
"body_vec".to_string(),
hnsw_config(4, DistanceMetric::Cosine),
);
let index =
MultiFieldVectorIndex::open_or_create(storage(), &fields, Arc::new(MockEmbedder)).unwrap();
let mut writer = index.writer().unwrap();
writer
.add_vectors(vec![
(1, "title_vec".to_string(), vec_of(&[1.0, 0.0, 0.0, 0.0])),
(1, "body_vec".to_string(), vec_of(&[0.0, 1.0, 0.0, 0.0])),
])
.unwrap();
writer.commit().unwrap();
let reader = index.reader().unwrap();
let title_vectors = reader.get_vectors_by_field("title_vec").unwrap();
let body_vectors = reader.get_vectors_by_field("body_vec").unwrap();
assert_eq!(
title_vectors.len(),
1,
"title_vec must keep its vector for doc 1"
);
assert_eq!(
body_vectors.len(),
1,
"body_vec must keep its vector for doc 1"
);
assert_eq!(title_vectors[0].0, 1);
assert_eq!(body_vectors[0].0, 1);
let stats = index.stats().unwrap();
assert_eq!(
stats.vector_count, 2,
"one doc in two fields = 2 vectors, not 1"
);
}
#[test]
fn field_routing_search_is_isolated() {
let mut fields = BTreeMap::new();
fields.insert(
"title_vec".to_string(),
hnsw_config(3, DistanceMetric::Cosine),
);
fields.insert(
"body_vec".to_string(),
hnsw_config(3, DistanceMetric::Cosine),
);
let index =
MultiFieldVectorIndex::open_or_create(storage(), &fields, Arc::new(MockEmbedder)).unwrap();
let mut writer = index.writer().unwrap();
writer
.add_vectors(vec![
(1, "title_vec".to_string(), vec_of(&[1.0, 0.0, 0.0])),
(2, "body_vec".to_string(), vec_of(&[1.0, 0.0, 0.0])),
])
.unwrap();
writer.commit().unwrap();
let searcher = index.searcher().unwrap();
let results = searcher
.search(&query(vec_of(&[1.0, 0.0, 0.0]), Some("title_vec"), 10))
.unwrap();
let ids: Vec<u64> = results.results.iter().map(|r| r.doc_id).collect();
assert_eq!(ids, vec![1], "title_vec query must only return doc 1");
}
#[test]
fn fieldless_query_prunes_dimension_mismatched_fields() {
let mut fields = BTreeMap::new();
fields.insert(
"small_vec".to_string(),
hnsw_config(3, DistanceMetric::Cosine),
);
fields.insert(
"big_vec".to_string(),
hnsw_config(8, DistanceMetric::Cosine),
);
let index =
MultiFieldVectorIndex::open_or_create(storage(), &fields, Arc::new(MockEmbedder)).unwrap();
let mut writer = index.writer().unwrap();
writer
.add_vectors(vec![
(1, "small_vec".to_string(), vec_of(&[1.0, 0.0, 0.0])),
(2, "big_vec".to_string(), vec_of(&[0.0; 8])),
])
.unwrap();
writer.commit().unwrap();
let searcher = index.searcher().unwrap();
let results = searcher
.search(&query(vec_of(&[1.0, 0.0, 0.0]), None, 10))
.unwrap();
let ids: Vec<u64> = results.results.iter().map(|r| r.doc_id).collect();
assert_eq!(
ids,
vec![1],
"only the dimension-matching field may contribute hits"
);
}
#[test]
fn fieldless_query_merges_across_homogeneous_fields() {
let mut fields = BTreeMap::new();
fields.insert("a_vec".to_string(), hnsw_config(3, DistanceMetric::Cosine));
fields.insert("b_vec".to_string(), hnsw_config(3, DistanceMetric::Cosine));
let index =
MultiFieldVectorIndex::open_or_create(storage(), &fields, Arc::new(MockEmbedder)).unwrap();
let mut writer = index.writer().unwrap();
writer
.add_vectors(vec![
(1, "a_vec".to_string(), vec_of(&[1.0, 0.0, 0.0])),
(2, "b_vec".to_string(), vec_of(&[1.0, 0.0, 0.0])),
])
.unwrap();
writer.commit().unwrap();
let searcher = index.searcher().unwrap();
let results = searcher
.search(&query(vec_of(&[1.0, 0.0, 0.0]), None, 10))
.unwrap();
let mut ids: Vec<u64> = results.results.iter().map(|r| r.doc_id).collect();
ids.sort_unstable();
assert_eq!(
ids,
vec![1, 2],
"both fields must contribute to a field-less query"
);
}
#[test]
fn unknown_field_rejects_whole_batch() {
let mut fields = BTreeMap::new();
fields.insert(
"title_vec".to_string(),
hnsw_config(3, DistanceMetric::Cosine),
);
let index =
MultiFieldVectorIndex::open_or_create(storage(), &fields, Arc::new(MockEmbedder)).unwrap();
let mut writer = index.writer().unwrap();
let err = writer
.add_vectors(vec![
(1, "title_vec".to_string(), vec_of(&[1.0, 0.0, 0.0])),
(1, "no_such_field".to_string(), vec_of(&[1.0, 0.0, 0.0])),
])
.unwrap_err();
assert!(
format!("{err:?}").contains("no_such_field"),
"error must name the unknown field: {err:?}"
);
writer.commit().unwrap();
let reader = index.reader().unwrap();
assert_eq!(reader.get_vectors_by_field("title_vec").unwrap().len(), 0);
}
#[test]
fn remove_field_does_not_affect_others() {
let mut fields = BTreeMap::new();
fields.insert(
"title_vec".to_string(),
hnsw_config(3, DistanceMetric::Cosine),
);
fields.insert(
"body_vec".to_string(),
hnsw_config(3, DistanceMetric::Cosine),
);
let index =
MultiFieldVectorIndex::open_or_create(storage(), &fields, Arc::new(MockEmbedder)).unwrap();
let mut writer = index.writer().unwrap();
writer
.add_vectors(vec![
(1, "title_vec".to_string(), vec_of(&[1.0, 0.0, 0.0])),
(2, "body_vec".to_string(), vec_of(&[0.0, 1.0, 0.0])),
])
.unwrap();
writer.commit().unwrap();
index.remove_field("title_vec").unwrap();
assert_eq!(
index.field_dimensions().keys().collect::<Vec<_>>(),
vec!["body_vec"],
"removed field must no longer be routed to"
);
let reader = index.reader().unwrap();
assert_eq!(reader.get_vectors_by_field("body_vec").unwrap().len(), 1);
}
#[test]
fn batch_search_preserves_query_order() {
let mut fields = BTreeMap::new();
fields.insert("a_vec".to_string(), hnsw_config(3, DistanceMetric::Cosine));
fields.insert("b_vec".to_string(), hnsw_config(3, DistanceMetric::Cosine));
let index =
MultiFieldVectorIndex::open_or_create(storage(), &fields, Arc::new(MockEmbedder)).unwrap();
let mut writer = index.writer().unwrap();
writer
.add_vectors(vec![
(1, "a_vec".to_string(), vec_of(&[1.0, 0.0, 0.0])),
(2, "b_vec".to_string(), vec_of(&[0.0, 1.0, 0.0])),
])
.unwrap();
writer.commit().unwrap();
let searcher = index.searcher().unwrap();
let queries = vec![
query(vec_of(&[0.0, 1.0, 0.0]), Some("b_vec"), 10), query(vec_of(&[1.0, 0.0, 0.0]), None, 10), query(vec_of(&[1.0, 0.0, 0.0]), Some("a_vec"), 10), ];
let results = searcher.search_batch(&queries).unwrap();
assert_eq!(results.len(), 3);
assert_eq!(results[0].results.first().map(|r| r.doc_id), Some(2));
assert_eq!(results[1].results.first().map(|r| r.doc_id), Some(1));
assert_eq!(results[2].results.first().map(|r| r.doc_id), Some(1));
}
#[test]
fn add_field_seeds_wal_seq_from_current_minimum() {
let mut fields = BTreeMap::new();
fields.insert(
"title_vec".to_string(),
hnsw_config(3, DistanceMetric::Cosine),
);
let index =
MultiFieldVectorIndex::open_or_create(storage(), &fields, Arc::new(MockEmbedder)).unwrap();
assert!(index.supports_dynamic_fields());
index.set_last_wal_seq(42).unwrap();
index.persist_deletions().unwrap();
assert_eq!(index.last_wal_seq(), 42);
index
.add_field("body_vec", hnsw_config(3, DistanceMetric::Cosine))
.unwrap();
index.persist_deletions().unwrap();
assert_eq!(index.last_wal_seq(), 42);
assert_eq!(
index.field_dimensions().keys().collect::<Vec<_>>(),
vec!["body_vec", "title_vec"]
);
assert!(
index
.add_field("body_vec", hnsw_config(3, DistanceMetric::Cosine))
.is_err()
);
}