use crate::error::{DbError, DbResult};
use crate::storage::index::{
VectorIndexConfig, VectorIndexStats, VectorMetric, VectorQuantization,
};
use rand::prelude::*;
use serde::{Deserialize, Serialize};
use std::cmp::Ordering;
use std::collections::{BinaryHeap, HashMap, HashSet};
use std::sync::RwLock;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct VectorSearchResult {
pub doc_key: String,
pub score: f32,
}
const DEFAULT_HNSW_THRESHOLD: usize = 10_000;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ScalarQuantParams {
pub min_vals: Vec<f32>,
pub max_vals: Vec<f32>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct QuantizedVectors {
vectors: HashMap<String, Vec<u8>>,
params: ScalarQuantParams,
}
impl QuantizedVectors {
pub fn from_full_vectors(vectors: &HashMap<String, Vec<f32>>, dimension: usize) -> Self {
if vectors.is_empty() {
return Self {
vectors: HashMap::new(),
params: ScalarQuantParams {
min_vals: vec![0.0; dimension],
max_vals: vec![1.0; dimension],
},
};
}
let mut min_vals = vec![f32::MAX; dimension];
let mut max_vals = vec![f32::MIN; dimension];
for vec in vectors.values() {
for (i, &v) in vec.iter().enumerate() {
min_vals[i] = min_vals[i].min(v);
max_vals[i] = max_vals[i].max(v);
}
}
let quantized: HashMap<String, Vec<u8>> = vectors
.iter()
.map(|(k, v)| (k.clone(), Self::quantize_vector(v, &min_vals, &max_vals)))
.collect();
Self {
vectors: quantized,
params: ScalarQuantParams { min_vals, max_vals },
}
}
fn quantize_vector(vec: &[f32], min_vals: &[f32], max_vals: &[f32]) -> Vec<u8> {
vec.iter()
.enumerate()
.map(|(i, &v)| {
let range = max_vals[i] - min_vals[i];
if range < 1e-10 {
127u8 } else {
((v - min_vals[i]) / range * 255.0).clamp(0.0, 255.0) as u8
}
})
.collect()
}
#[allow(dead_code)]
pub fn dequantize_vector(&self, quantized: &[u8]) -> Vec<f32> {
quantized
.iter()
.enumerate()
.map(|(i, &q)| {
let range = self.params.max_vals[i] - self.params.min_vals[i];
self.params.min_vals[i] + (q as f32 / 255.0) * range
})
.collect()
}
pub fn insert(&mut self, doc_key: &str, vec: &[f32]) {
let quantized = Self::quantize_vector(vec, &self.params.min_vals, &self.params.max_vals);
self.vectors.insert(doc_key.to_string(), quantized);
}
pub fn remove(&mut self, doc_key: &str) {
self.vectors.remove(doc_key);
}
pub fn len(&self) -> usize {
self.vectors.len()
}
#[allow(dead_code)]
pub fn is_empty(&self) -> bool {
self.vectors.is_empty()
}
pub fn params(&self) -> &ScalarQuantParams {
&self.params
}
pub fn vectors(&self) -> &HashMap<String, Vec<u8>> {
&self.vectors
}
pub fn clear(&mut self) {
self.vectors.clear();
}
}
pub fn cosine_similarity_asymmetric(
query: &[f32],
quantized: &[u8],
params: &ScalarQuantParams,
) -> f32 {
let query_norm_sq: f32 = query.iter().map(|x| x * x).sum();
if query_norm_sq < 1e-10 {
return 0.0;
}
let mut dot = 0.0f32;
let mut quantized_norm_sq = 0.0f32;
for i in 0..query.len() {
let range = params.max_vals[i] - params.min_vals[i];
let dequant = params.min_vals[i] + (quantized[i] as f32 / 255.0) * range;
dot += query[i] * dequant;
quantized_norm_sq += dequant * dequant;
}
let query_norm = query_norm_sq.sqrt();
let quantized_norm = quantized_norm_sq.sqrt().max(1e-10);
dot / (query_norm * quantized_norm)
}
pub fn euclidean_distance_asymmetric(
query: &[f32],
quantized: &[u8],
params: &ScalarQuantParams,
) -> f32 {
let mut sum_sq = 0.0f32;
for i in 0..query.len() {
let range = params.max_vals[i] - params.min_vals[i];
let dequant = params.min_vals[i] + (quantized[i] as f32 / 255.0) * range;
let diff = query[i] - dequant;
sum_sq += diff * diff;
}
sum_sq.sqrt()
}
pub fn dot_product_asymmetric(query: &[f32], quantized: &[u8], params: &ScalarQuantParams) -> f32 {
let mut dot = 0.0f32;
for i in 0..query.len() {
let range = params.max_vals[i] - params.min_vals[i];
let dequant = params.min_vals[i] + (quantized[i] as f32 / 255.0) * range;
dot += query[i] * dequant;
}
dot
}
#[derive(Debug, Clone, Serialize, Deserialize)]
struct HnswNode {
neighbors: Vec<Vec<String>>,
level: usize,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
struct HnswGraph {
nodes: HashMap<String, HnswNode>,
entry_point: Option<String>,
max_level: usize,
m: usize,
m0: usize,
ef_construction: usize,
level_mult: f64,
}
#[derive(Clone)]
struct Neighbor {
doc_key: String,
distance: f32,
}
impl PartialEq for Neighbor {
fn eq(&self, other: &Self) -> bool {
self.distance == other.distance
}
}
impl Eq for Neighbor {}
impl PartialOrd for Neighbor {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
Some(self.cmp(other))
}
}
impl Ord for Neighbor {
fn cmp(&self, other: &Self) -> Ordering {
other
.distance
.partial_cmp(&self.distance)
.unwrap_or(Ordering::Equal)
}
}
#[derive(Clone)]
struct MaxNeighbor(Neighbor);
impl PartialEq for MaxNeighbor {
fn eq(&self, other: &Self) -> bool {
self.0.distance == other.0.distance
}
}
impl Eq for MaxNeighbor {}
impl PartialOrd for MaxNeighbor {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
Some(self.cmp(other))
}
}
impl Ord for MaxNeighbor {
fn cmp(&self, other: &Self) -> Ordering {
self.0
.distance
.partial_cmp(&other.0.distance)
.unwrap_or(Ordering::Equal)
}
}
impl HnswGraph {
fn new(m: usize, ef_construction: usize) -> Self {
let m = m.max(2); Self {
nodes: HashMap::new(),
entry_point: None,
max_level: 0,
m,
m0: m * 2, ef_construction,
level_mult: 1.0 / (m as f64).ln(),
}
}
fn random_level(&self) -> usize {
let mut rng = rand::thread_rng();
let r: f64 = rng.gen();
let level = (-r.ln() * self.level_mult).floor() as usize;
level.min(16)
}
fn calculate_distance(
query: &[f32],
doc_key: &str,
vectors: &HashMap<String, Vec<f32>>,
metric: VectorMetric,
) -> f32 {
match vectors.get(doc_key) {
Some(vec) => match metric {
VectorMetric::Cosine => 1.0 - cosine_similarity(query, vec),
VectorMetric::Euclidean => euclidean_distance(query, vec),
VectorMetric::DotProduct => -dot_product(query, vec), },
None => f32::MAX,
}
}
fn search_layer(
&self,
query: &[f32],
entry_points: &[String],
ef: usize,
level: usize,
vectors: &HashMap<String, Vec<f32>>,
metric: VectorMetric,
) -> Vec<Neighbor> {
let mut visited: HashSet<String> = HashSet::new();
let mut candidates: BinaryHeap<Neighbor> = BinaryHeap::new();
let mut result: BinaryHeap<MaxNeighbor> = BinaryHeap::new();
for ep in entry_points {
if visited.insert(ep.clone()) {
let dist = Self::calculate_distance(query, ep, vectors, metric);
let neighbor = Neighbor {
doc_key: ep.clone(),
distance: dist,
};
candidates.push(neighbor.clone());
result.push(MaxNeighbor(neighbor));
}
}
while let Some(current) = candidates.pop() {
let furthest_dist = result.peek().map(|n| n.0.distance).unwrap_or(f32::MAX);
if current.distance > furthest_dist {
break;
}
if let Some(node) = self.nodes.get(¤t.doc_key) {
if level < node.neighbors.len() {
for neighbor_key in &node.neighbors[level] {
if visited.insert(neighbor_key.clone()) {
let dist =
Self::calculate_distance(query, neighbor_key, vectors, metric);
let furthest_dist =
result.peek().map(|n| n.0.distance).unwrap_or(f32::MAX);
if dist < furthest_dist || result.len() < ef {
let neighbor = Neighbor {
doc_key: neighbor_key.clone(),
distance: dist,
};
candidates.push(neighbor.clone());
result.push(MaxNeighbor(neighbor));
while result.len() > ef {
result.pop();
}
}
}
}
}
}
}
let mut results: Vec<Neighbor> = result.into_iter().map(|mn| mn.0).collect();
results.sort_by(|a, b| {
a.distance
.partial_cmp(&b.distance)
.unwrap_or(Ordering::Equal)
});
results
}
fn select_neighbors(&self, candidates: &[Neighbor], m: usize) -> Vec<String> {
candidates
.iter()
.take(m)
.map(|n| n.doc_key.clone())
.collect()
}
fn insert(
&mut self,
doc_key: &str,
_vector: &[f32],
vectors: &HashMap<String, Vec<f32>>,
metric: VectorMetric,
) {
if self.nodes.contains_key(doc_key) {
self.remove(doc_key);
}
let query = match vectors.get(doc_key) {
Some(v) => v,
None => return, };
let node_level = self.random_level();
let mut new_node = HnswNode {
neighbors: vec![Vec::new(); node_level + 1],
level: node_level,
};
if self.entry_point.is_none() {
self.entry_point = Some(doc_key.to_string());
self.max_level = node_level;
self.nodes.insert(doc_key.to_string(), new_node);
return;
}
let entry_point = self.entry_point.clone().unwrap();
let mut current_ep = vec![entry_point];
for level in (node_level + 1..=self.max_level).rev() {
let nearest = self.search_layer(query, ¤t_ep, 1, level, vectors, metric);
if !nearest.is_empty() {
current_ep = vec![nearest[0].doc_key.clone()];
}
}
for level in (0..=node_level.min(self.max_level)).rev() {
let m_level = if level == 0 { self.m0 } else { self.m };
let candidates = self.search_layer(
query,
¤t_ep,
self.ef_construction,
level,
vectors,
metric,
);
let selected = self.select_neighbors(&candidates, m_level);
if level < new_node.neighbors.len() {
new_node.neighbors[level] = selected.clone();
}
for neighbor_key in &selected {
if let Some(neighbor_node) = self.nodes.get_mut(neighbor_key) {
if level < neighbor_node.neighbors.len() {
neighbor_node.neighbors[level].push(doc_key.to_string());
if neighbor_node.neighbors[level].len() > m_level {
let mut neighbor_candidates: Vec<Neighbor> = neighbor_node.neighbors
[level]
.iter()
.filter_map(|k| {
vectors.get(neighbor_key).map(|nv| Neighbor {
doc_key: k.clone(),
distance: Self::calculate_distance(nv, k, vectors, metric),
})
})
.collect();
neighbor_candidates.sort_by(|a, b| {
a.distance
.partial_cmp(&b.distance)
.unwrap_or(Ordering::Equal)
});
neighbor_node.neighbors[level] = neighbor_candidates
.iter()
.take(m_level)
.map(|n| n.doc_key.clone())
.collect();
}
}
}
}
if !candidates.is_empty() {
current_ep = candidates.iter().map(|n| n.doc_key.clone()).collect();
}
}
self.nodes.insert(doc_key.to_string(), new_node);
if node_level > self.max_level {
self.entry_point = Some(doc_key.to_string());
self.max_level = node_level;
}
}
fn remove(&mut self, doc_key: &str) {
if let Some(removed_node) = self.nodes.remove(doc_key) {
for level in 0..removed_node.neighbors.len() {
for neighbor_key in &removed_node.neighbors[level] {
if let Some(neighbor_node) = self.nodes.get_mut(neighbor_key) {
if level < neighbor_node.neighbors.len() {
neighbor_node.neighbors[level].retain(|k| k != doc_key);
}
}
}
}
if self.entry_point.as_ref() == Some(&doc_key.to_string()) {
self.entry_point = self.nodes.keys().next().cloned();
if let Some(ep) = &self.entry_point {
if let Some(node) = self.nodes.get(ep) {
self.max_level = node.level;
}
} else {
self.max_level = 0;
}
}
}
}
fn search(
&self,
query: &[f32],
k: usize,
ef: usize,
vectors: &HashMap<String, Vec<f32>>,
metric: VectorMetric,
) -> Vec<Neighbor> {
if self.entry_point.is_none() {
return vec![];
}
let entry_point = self.entry_point.clone().unwrap();
let mut current_ep = vec![entry_point];
for level in (1..=self.max_level).rev() {
let nearest = self.search_layer(query, ¤t_ep, 1, level, vectors, metric);
if !nearest.is_empty() {
current_ep = vec![nearest[0].doc_key.clone()];
}
}
let ef_search = ef.max(k);
let mut results = self.search_layer(query, ¤t_ep, ef_search, 0, vectors, metric);
results.truncate(k);
results
}
fn clear(&mut self) {
self.nodes.clear();
self.entry_point = None;
self.max_level = 0;
}
}
pub struct VectorIndex {
config: VectorIndexConfig,
vectors: RwLock<HashMap<String, Vec<f32>>>,
quantized_vectors: RwLock<Option<QuantizedVectors>>,
hnsw_graph: RwLock<Option<HnswGraph>>,
}
impl VectorIndex {
pub fn new(config: VectorIndexConfig) -> DbResult<Self> {
if config.dimension == 0 {
return Err(DbError::BadRequest(
"Vector dimension must be greater than 0".to_string(),
));
}
Ok(Self {
config,
vectors: RwLock::new(HashMap::new()),
quantized_vectors: RwLock::new(None),
hnsw_graph: RwLock::new(None),
})
}
fn hnsw_threshold(&self) -> usize {
DEFAULT_HNSW_THRESHOLD
}
fn build_hnsw_graph(&self) {
let vectors = self.vectors.read().unwrap();
if vectors.len() < self.hnsw_threshold() {
return;
}
let mut graph = HnswGraph::new(self.config.m, self.config.ef_construction);
for (doc_key, vector) in vectors.iter() {
graph.insert(doc_key, vector, &vectors, self.config.metric);
}
let mut hnsw = self.hnsw_graph.write().unwrap();
*hnsw = Some(graph);
}
pub fn config(&self) -> &VectorIndexConfig {
&self.config
}
pub fn len(&self) -> usize {
self.vectors.read().unwrap().len()
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn clear(&self) {
let mut vectors = self.vectors.write().unwrap();
vectors.clear();
let mut quantized = self.quantized_vectors.write().unwrap();
if let Some(q) = quantized.as_mut() {
q.clear();
}
*quantized = None;
let mut hnsw = self.hnsw_graph.write().unwrap();
if let Some(graph) = hnsw.as_mut() {
graph.clear();
}
*hnsw = None;
}
pub fn insert(&self, doc_key: &str, vector: &[f32]) -> DbResult<()> {
if vector.len() != self.config.dimension {
return Err(DbError::BadRequest(format!(
"Vector dimension mismatch: expected {}, got {}",
self.config.dimension,
vector.len()
)));
}
{
let mut vectors = self.vectors.write().unwrap();
vectors.insert(doc_key.to_string(), vector.to_vec());
}
{
let mut quantized = self.quantized_vectors.write().unwrap();
if let Some(q) = quantized.as_mut() {
q.insert(doc_key, vector);
}
}
let vectors = self.vectors.read().unwrap();
let count = vectors.len();
drop(vectors);
let mut hnsw = self.hnsw_graph.write().unwrap();
if let Some(graph) = hnsw.as_mut() {
let vectors = self.vectors.read().unwrap();
graph.insert(doc_key, vector, &vectors, self.config.metric);
} else if count >= self.hnsw_threshold() {
drop(hnsw);
self.build_hnsw_graph();
}
Ok(())
}
pub fn remove(&self, doc_key: &str) -> DbResult<bool> {
let removed = {
let mut vectors = self.vectors.write().unwrap();
vectors.remove(doc_key).is_some()
};
if removed {
{
let mut quantized = self.quantized_vectors.write().unwrap();
if let Some(q) = quantized.as_mut() {
q.remove(doc_key);
}
}
let mut hnsw = self.hnsw_graph.write().unwrap();
if let Some(graph) = hnsw.as_mut() {
graph.remove(doc_key);
}
}
Ok(removed)
}
pub fn contains(&self, doc_key: &str) -> bool {
self.vectors.read().unwrap().contains_key(doc_key)
}
pub fn get(&self, doc_key: &str) -> Option<Vec<f32>> {
self.vectors.read().unwrap().get(doc_key).cloned()
}
pub fn search(
&self,
query: &[f32],
limit: usize,
ef: usize,
) -> DbResult<Vec<VectorSearchResult>> {
if query.len() != self.config.dimension {
return Err(DbError::BadRequest(format!(
"Query vector dimension mismatch: expected {}, got {}",
self.config.dimension,
query.len()
)));
}
{
let hnsw = self.hnsw_graph.read().unwrap();
if let Some(graph) = hnsw.as_ref() {
let vectors = self.vectors.read().unwrap();
let ef_search = ef.max(limit * 2).max(40);
let hnsw_results =
graph.search(query, limit, ef_search, &vectors, self.config.metric);
let results: Vec<VectorSearchResult> = hnsw_results
.into_iter()
.map(|n| {
let score = match self.config.metric {
VectorMetric::Cosine => 1.0 - n.distance, VectorMetric::Euclidean => n.distance, VectorMetric::DotProduct => -n.distance, };
VectorSearchResult {
doc_key: n.doc_key,
score,
}
})
.collect();
return Ok(results);
}
}
self.brute_force_search(query, limit)
}
fn brute_force_search(&self, query: &[f32], limit: usize) -> DbResult<Vec<VectorSearchResult>> {
{
let quantized = self.quantized_vectors.read().unwrap();
if let Some(q) = quantized.as_ref() {
return self.brute_force_search_quantized(query, limit, q);
}
}
let vectors = self.vectors.read().unwrap();
let mut results: Vec<VectorSearchResult> = vectors
.iter()
.map(|(doc_key, vec)| {
let score = match self.config.metric {
VectorMetric::Cosine => cosine_similarity(query, vec),
VectorMetric::Euclidean => euclidean_distance(query, vec),
VectorMetric::DotProduct => dot_product(query, vec),
};
VectorSearchResult {
doc_key: doc_key.clone(),
score,
}
})
.collect();
match self.config.metric {
VectorMetric::Cosine | VectorMetric::DotProduct => {
results.sort_by(|a, b| b.score.partial_cmp(&a.score).unwrap_or(Ordering::Equal));
}
VectorMetric::Euclidean => {
results.sort_by(|a, b| a.score.partial_cmp(&b.score).unwrap_or(Ordering::Equal));
}
}
results.truncate(limit);
Ok(results)
}
fn brute_force_search_quantized(
&self,
query: &[f32],
limit: usize,
quantized: &QuantizedVectors,
) -> DbResult<Vec<VectorSearchResult>> {
let params = quantized.params();
let mut results: Vec<VectorSearchResult> = quantized
.vectors()
.iter()
.map(|(doc_key, quantized_vec)| {
let score = match self.config.metric {
VectorMetric::Cosine => {
cosine_similarity_asymmetric(query, quantized_vec, params)
}
VectorMetric::Euclidean => {
euclidean_distance_asymmetric(query, quantized_vec, params)
}
VectorMetric::DotProduct => {
dot_product_asymmetric(query, quantized_vec, params)
}
};
VectorSearchResult {
doc_key: doc_key.clone(),
score,
}
})
.collect();
match self.config.metric {
VectorMetric::Cosine | VectorMetric::DotProduct => {
results.sort_by(|a, b| b.score.partial_cmp(&a.score).unwrap_or(Ordering::Equal));
}
VectorMetric::Euclidean => {
results.sort_by(|a, b| a.score.partial_cmp(&b.score).unwrap_or(Ordering::Equal));
}
}
results.truncate(limit);
Ok(results)
}
pub fn similarity(&self, doc_key: &str, query: &[f32]) -> DbResult<Option<f32>> {
if query.len() != self.config.dimension {
return Err(DbError::BadRequest(format!(
"Query vector dimension mismatch: expected {}, got {}",
self.config.dimension,
query.len()
)));
}
let vectors = self.vectors.read().unwrap();
let doc_vec = match vectors.get(doc_key) {
Some(v) => v,
None => return Ok(None),
};
let score = match self.config.metric {
VectorMetric::Cosine => cosine_similarity(query, doc_vec),
VectorMetric::Euclidean => euclidean_distance(query, doc_vec),
VectorMetric::DotProduct => dot_product(query, doc_vec),
};
Ok(Some(score))
}
pub fn serialize(&self) -> DbResult<Vec<u8>> {
let vectors = self.vectors.read().unwrap();
let quantized = self.quantized_vectors.read().unwrap();
let hnsw = self.hnsw_graph.read().unwrap();
let data = VectorIndexDataV3 {
config: self.config.clone(),
vectors: vectors.clone(),
quantized_vectors: quantized.clone(),
hnsw_graph: hnsw.clone(),
};
bincode::serialize(&data)
.map_err(|e| DbError::InternalError(format!("Serialization error: {}", e)))
}
pub fn deserialize(bytes: &[u8]) -> DbResult<Self> {
if let Ok(data) = bincode::deserialize::<VectorIndexDataV3>(bytes) {
return Ok(Self {
config: data.config,
vectors: RwLock::new(data.vectors),
quantized_vectors: RwLock::new(data.quantized_vectors),
hnsw_graph: RwLock::new(data.hnsw_graph),
});
}
if let Ok(data) = bincode::deserialize::<VectorIndexDataV2>(bytes) {
return Ok(Self {
config: data.config,
vectors: RwLock::new(data.vectors),
quantized_vectors: RwLock::new(None),
hnsw_graph: RwLock::new(data.hnsw_graph),
});
}
let data: VectorIndexData = bincode::deserialize(bytes)
.map_err(|e| DbError::InternalError(format!("Deserialization error: {}", e)))?;
Ok(Self {
config: data.config,
vectors: RwLock::new(data.vectors),
quantized_vectors: RwLock::new(None),
hnsw_graph: RwLock::new(None),
})
}
pub fn is_hnsw_active(&self) -> bool {
self.hnsw_graph.read().unwrap().is_some()
}
pub fn is_quantized(&self) -> bool {
self.quantized_vectors.read().unwrap().is_some()
}
pub fn quantization_type(&self) -> VectorQuantization {
if self.quantized_vectors.read().unwrap().is_some() {
VectorQuantization::Scalar
} else {
VectorQuantization::None
}
}
pub fn quantize(&self) -> DbResult<usize> {
let vectors = self.vectors.read().unwrap();
let count = vectors.len();
if count == 0 {
return Ok(0);
}
let quantized = QuantizedVectors::from_full_vectors(&vectors, self.config.dimension);
drop(vectors);
let mut quantized_lock = self.quantized_vectors.write().unwrap();
*quantized_lock = Some(quantized);
Ok(count)
}
pub fn dequantize(&self) {
let mut quantized = self.quantized_vectors.write().unwrap();
*quantized = None;
}
pub fn quantization_stats(&self) -> Option<QuantizationStats> {
let quantized = self.quantized_vectors.read().unwrap();
quantized.as_ref().map(|q| {
let vector_count = q.len();
let memory_bytes = vector_count * self.config.dimension; let full_memory_bytes =
vector_count * self.config.dimension * std::mem::size_of::<f32>();
QuantizationStats {
vector_count,
memory_bytes,
full_memory_bytes,
compression_ratio: if memory_bytes > 0 {
full_memory_bytes as f32 / memory_bytes as f32
} else {
1.0
},
}
})
}
pub fn stats(&self) -> VectorIndexStats {
let (memory_bytes, compression_ratio) = if let Some(stats) = self.quantization_stats() {
(stats.memory_bytes, stats.compression_ratio)
} else {
(0, 1.0)
};
VectorIndexStats {
name: self.config.name.clone(),
field: self.config.field.clone(),
dimension: self.config.dimension,
metric: self.config.metric,
m: self.config.m,
ef_construction: self.config.ef_construction,
indexed_vectors: self.len(),
quantization: self.quantization_type(),
memory_bytes,
compression_ratio,
}
}
}
#[derive(Debug, Clone, Serialize)]
pub struct QuantizationStats {
pub vector_count: usize,
pub memory_bytes: usize,
pub full_memory_bytes: usize,
pub compression_ratio: f32,
}
#[derive(Serialize, Deserialize)]
struct VectorIndexData {
config: VectorIndexConfig,
vectors: HashMap<String, Vec<f32>>,
}
#[derive(Serialize, Deserialize)]
struct VectorIndexDataV2 {
config: VectorIndexConfig,
vectors: HashMap<String, Vec<f32>>,
#[serde(default)]
hnsw_graph: Option<HnswGraph>,
}
#[derive(Serialize, Deserialize)]
struct VectorIndexDataV3 {
config: VectorIndexConfig,
vectors: HashMap<String, Vec<f32>>,
#[serde(default)]
quantized_vectors: Option<QuantizedVectors>,
#[serde(default)]
hnsw_graph: Option<HnswGraph>,
}
pub fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
if a.len() != b.len() || a.is_empty() {
return 0.0;
}
let mut dot = 0.0f32;
let mut norm_a = 0.0f32;
let mut norm_b = 0.0f32;
for i in 0..a.len() {
dot += a[i] * b[i];
norm_a += a[i] * a[i];
norm_b += b[i] * b[i];
}
let norm = (norm_a.sqrt() * norm_b.sqrt()).max(1e-10);
dot / norm
}
pub fn euclidean_distance(a: &[f32], b: &[f32]) -> f32 {
if a.len() != b.len() {
return f32::MAX;
}
let mut sum = 0.0f32;
for i in 0..a.len() {
let diff = a[i] - b[i];
sum += diff * diff;
}
sum.sqrt()
}
pub fn dot_product(a: &[f32], b: &[f32]) -> f32 {
if a.len() != b.len() {
return 0.0;
}
let mut dot = 0.0f32;
for i in 0..a.len() {
dot += a[i] * b[i];
}
dot
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_cosine_similarity() {
let a = vec![1.0, 0.0, 0.0];
let b = vec![1.0, 0.0, 0.0];
assert!((cosine_similarity(&a, &b) - 1.0).abs() < 1e-6);
let c = vec![0.0, 1.0, 0.0];
assert!(cosine_similarity(&a, &c).abs() < 1e-6);
let d = vec![-1.0, 0.0, 0.0];
assert!((cosine_similarity(&a, &d) + 1.0).abs() < 1e-6);
}
#[test]
fn test_euclidean_distance() {
let a = vec![0.0, 0.0, 0.0];
let b = vec![1.0, 0.0, 0.0];
assert!((euclidean_distance(&a, &b) - 1.0).abs() < 1e-6);
let c = vec![3.0, 4.0, 0.0];
assert!((euclidean_distance(&a, &c) - 5.0).abs() < 1e-6);
}
#[test]
fn test_dot_product() {
let a = vec![1.0, 2.0, 3.0];
let b = vec![4.0, 5.0, 6.0];
assert!((dot_product(&a, &b) - 32.0).abs() < 1e-6);
}
#[test]
fn test_vector_index_creation() {
let config = VectorIndexConfig::new("test_idx".to_string(), "embedding".to_string(), 3);
let index = VectorIndex::new(config).unwrap();
assert!(index.is_empty());
}
#[test]
fn test_vector_index_insert_and_search() {
let config = VectorIndexConfig::new("test_idx".to_string(), "embedding".to_string(), 3);
let index = VectorIndex::new(config).unwrap();
index.insert("doc1", &[1.0, 0.0, 0.0]).unwrap();
index.insert("doc2", &[0.0, 1.0, 0.0]).unwrap();
index.insert("doc3", &[0.9, 0.1, 0.0]).unwrap();
assert_eq!(index.len(), 3);
let results = index.search(&[1.0, 0.0, 0.0], 2, 10).unwrap();
assert_eq!(results.len(), 2);
assert_eq!(results[0].doc_key, "doc1");
assert!((results[0].score - 1.0).abs() < 1e-6);
}
#[test]
fn test_vector_index_dimension_validation() {
let config = VectorIndexConfig::new("test_idx".to_string(), "embedding".to_string(), 3);
let index = VectorIndex::new(config).unwrap();
let result = index.insert("doc1", &[1.0, 0.0]);
assert!(result.is_err());
}
#[test]
fn test_vector_index_remove() {
let config = VectorIndexConfig::new("test_idx".to_string(), "embedding".to_string(), 3);
let index = VectorIndex::new(config).unwrap();
index.insert("doc1", &[1.0, 0.0, 0.0]).unwrap();
assert_eq!(index.len(), 1);
let removed = index.remove("doc1").unwrap();
assert!(removed);
assert_eq!(index.len(), 0);
let removed = index.remove("nonexistent").unwrap();
assert!(!removed);
}
#[test]
fn test_vector_index_with_euclidean() {
let config = VectorIndexConfig::new("test_idx".to_string(), "embedding".to_string(), 3)
.with_metric(VectorMetric::Euclidean);
let index = VectorIndex::new(config).unwrap();
index.insert("doc1", &[0.0, 0.0, 0.0]).unwrap();
index.insert("doc2", &[1.0, 0.0, 0.0]).unwrap();
index.insert("doc3", &[10.0, 0.0, 0.0]).unwrap();
let results = index.search(&[0.0, 0.0, 0.0], 3, 10).unwrap();
assert_eq!(results[0].doc_key, "doc1");
assert!(results[0].score < 0.1);
}
#[test]
fn test_vector_index_similarity() {
let config = VectorIndexConfig::new("test_idx".to_string(), "embedding".to_string(), 3);
let index = VectorIndex::new(config).unwrap();
index.insert("doc1", &[1.0, 0.0, 0.0]).unwrap();
let sim = index.similarity("doc1", &[1.0, 0.0, 0.0]).unwrap();
assert!(sim.is_some());
assert!((sim.unwrap() - 1.0).abs() < 1e-6);
let sim = index.similarity("nonexistent", &[1.0, 0.0, 0.0]).unwrap();
assert!(sim.is_none());
}
#[test]
fn test_vector_index_serialize_deserialize() {
let config = VectorIndexConfig::new("test_idx".to_string(), "embedding".to_string(), 3);
let index = VectorIndex::new(config).unwrap();
index.insert("doc1", &[1.0, 0.0, 0.0]).unwrap();
index.insert("doc2", &[0.0, 1.0, 0.0]).unwrap();
let bytes = index.serialize().unwrap();
let restored = VectorIndex::deserialize(&bytes).unwrap();
assert_eq!(restored.len(), 2);
assert!(restored.contains("doc1"));
assert!(restored.contains("doc2"));
}
#[test]
fn test_vector_index_update() {
let config = VectorIndexConfig::new("test_idx".to_string(), "embedding".to_string(), 3);
let index = VectorIndex::new(config).unwrap();
index.insert("doc1", &[1.0, 0.0, 0.0]).unwrap();
index.insert("doc1", &[0.0, 1.0, 0.0]).unwrap();
assert_eq!(index.len(), 1);
let vec = index.get("doc1").unwrap();
assert!((vec[0] - 0.0).abs() < 1e-6);
assert!((vec[1] - 1.0).abs() < 1e-6);
}
}