diff --git a/Cargo.lock b/Cargo.lock index f46657a9..4e787c55 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -420,6 +420,17 @@ dependencies = [ "cc", ] +[[package]] +name = "codespan-reporting" +version = "0.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "af491d569909a7e4dee0ad7db7f5341fef5c614d5b8ec8cf765732aba3cff681" +dependencies = [ + "serde", + "termcolor", + "unicode-width", +] + [[package]] name = "colorchoice" version = "1.0.4" @@ -558,6 +569,68 @@ dependencies = [ "cipher", ] +[[package]] +name = "cxx" +version = "1.0.194" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "747d8437319e3a2f43d93b341c137927ca70c0f5dabeea7a005a73665e247c7e" +dependencies = [ + "cc", + "cxx-build", + "cxxbridge-cmd", + "cxxbridge-flags", + "cxxbridge-macro", + "foldhash 0.2.0", + "link-cplusplus", +] + +[[package]] +name = "cxx-build" +version = "1.0.194" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b0f4697d190a142477b16aef7da8a99bfdc41e7e8b1687583c0d23a79c7afc1e" +dependencies = [ + "cc", + "codespan-reporting", + "indexmap", + "proc-macro2", + "quote", + "scratch", + "syn 2.0.114", +] + +[[package]] +name = "cxxbridge-cmd" +version = "1.0.194" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d0956799fa8678d4c50eed028f2de1c0552ae183c76e976cf7ca8c4e36a7c328" +dependencies = [ + "clap", + "codespan-reporting", + "indexmap", + "proc-macro2", + "quote", + "syn 2.0.114", +] + +[[package]] +name = "cxxbridge-flags" +version = "1.0.194" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "23384a836ab4f0ad98ace7e3955ad2de39de42378ab487dc28d3990392cb283a" + +[[package]] +name = "cxxbridge-macro" +version = "1.0.194" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e6acc6b5822b9526adfb4fc377b67128fdd60aac757cc4a741a6278603f763cf" +dependencies = [ + "indexmap", + "proc-macro2", + "quote", + "syn 2.0.114", +] + [[package]] name = "dashmap" version = "6.1.0" @@ -743,6 +816,7 @@ dependencies = [ "thiserror 2.0.18", "tokio", "tracing", + "usearch", ] [[package]] @@ -808,6 +882,12 @@ version = "0.1.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d9c4f5dac5e15c24eb999c26181a6ca40b39fe946cbe4c263c7209467bc83af2" +[[package]] +name = "foldhash" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "77ce24cb58228fbb8aa041425bb1050850ac19177686ea6e0f41a70416f56fdb" + [[package]] name = "fs_extra" version = "1.3.0" @@ -1003,7 +1083,7 @@ version = "0.15.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9229cfe53dfd69f0609a49f65461bd93001ea1ef889cd5529dd176593f5338a1" dependencies = [ - "foldhash", + "foldhash 0.1.5", ] [[package]] @@ -1261,6 +1341,15 @@ dependencies = [ "libc", ] +[[package]] +name = "link-cplusplus" +version = "1.0.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7f78c730aaa7d0b9336a299029ea49f9ee53b0ed06e9202e8cb7db9bae7b8c82" +dependencies = [ + "cc", +] + [[package]] name = "linux-raw-sys" version = "0.11.0" @@ -2085,6 +2174,12 @@ version = "1.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "94143f37725109f92c262ed2cf5e59bce7498c01bcc1502d7b9afe439a4e9f49" +[[package]] +name = "scratch" +version = "1.0.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d68f2ec51b097e4c1a75b681a8bec621909b5e91f15bb7b840c4f2f7b01148b2" + [[package]] name = "seahash" version = "4.1.0" @@ -2275,6 +2370,15 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "termcolor" +version = "1.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "06794f8f6c5c898b3275aebefa6b8a1cb24cd2c6c79397ab15774837a0bc5755" +dependencies = [ + "winapi-util", +] + [[package]] name = "thiserror" version = "1.0.69" @@ -2579,6 +2683,16 @@ version = "0.9.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8ecb6da28b8a351d773b68d5825ac39017e680750f980f3a1a85cd8dd28a47c1" +[[package]] +name = "usearch" +version = "2.23.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0a03c05af8d678ec19f014c734ab667c20ea54128b4f9a1472cb470246a9b341" +dependencies = [ + "cxx", + "cxx-build", +] + [[package]] name = "utf8-width" version = "0.1.8" diff --git a/Cargo.toml b/Cargo.toml index d9085843..9357a627 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -64,6 +64,9 @@ rustls-native-certs = "0.8" # dynamic protobuf messages prost-reflect = "0.16" +# HNSW vector similarity search +usearch = "2.23" + # internal crates (version required for crates.io publishing) emberkv-core = { version = "0.4.3", path = "crates/ember-core" } ember-protocol = { version = "0.4.3", path = "crates/ember-protocol" } diff --git a/crates/ember-core/Cargo.toml b/crates/ember-core/Cargo.toml index c91b5f1c..b84a020c 100644 --- a/crates/ember-core/Cargo.toml +++ b/crates/ember-core/Cargo.toml @@ -12,6 +12,7 @@ readme = "README.md" [features] encryption = ["ember-persistence/encryption"] protobuf = ["prost-reflect", "ember-persistence/protobuf"] +vector = ["usearch", "ember-persistence/vector"] [lib] name = "ember_core" @@ -27,6 +28,7 @@ tracing = { workspace = true } rand = { workspace = true } ordered-float = { workspace = true } prost-reflect = { workspace = true, optional = true } +usearch = { workspace = true, optional = true } dashmap = "6" parking_lot = "0.12" diff --git a/crates/ember-core/src/keyspace.rs b/crates/ember-core/src/keyspace.rs index 90c52630..cf6fedd0 100644 --- a/crates/ember-core/src/keyspace.rs +++ b/crates/ember-core/src/keyspace.rs @@ -144,6 +144,30 @@ pub struct ZAddResult { pub applied: Vec<(f64, String)>, } +/// Result of a VADD operation, carrying the applied element for AOF persistence. +#[cfg(feature = "vector")] +#[derive(Debug, Clone)] +pub struct VAddResult { + /// The element name that was added or updated. + pub element: String, + /// The vector that was stored. + pub vector: Vec, + /// Whether a new element was added (false = updated existing). + pub added: bool, +} + +/// Errors from vector write operations. +#[cfg(feature = "vector")] +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum VectorWriteError { + /// The key holds a different type than expected. + WrongType, + /// Memory limit reached. + OutOfMemory, + /// usearch index error (dimension mismatch, capacity, etc). + IndexError(String), +} + /// How the keyspace should handle writes when the memory limit is reached. #[derive(Debug, Clone, Copy, Default, PartialEq, Eq)] pub enum EvictionPolicy { @@ -1850,6 +1874,218 @@ impl Keyspace { expired } + // -- vector operations -- + + /// Adds a vector to a vector set, creating the set if it doesn't exist. + /// + /// On first insert, the set's configuration (dim, metric, quantization, + /// connectivity, expansion_add) is locked. Subsequent inserts must match + /// the established dimensionality. + /// + /// Returns a `VAddResult` with the element name, vector, and whether it + /// was newly added (for AOF recording). + #[cfg(feature = "vector")] + #[allow(clippy::too_many_arguments)] + pub fn vadd( + &mut self, + key: &str, + element: String, + vector: Vec, + metric: crate::types::vector::DistanceMetric, + quantization: crate::types::vector::QuantizationType, + connectivity: usize, + expansion_add: usize, + ) -> Result { + use crate::types::vector::VectorSet; + + self.remove_if_expired(key); + + let is_new = !self.entries.contains_key(key); + if !is_new && !matches!(self.entries[key].value, Value::Vector(_)) { + return Err(VectorWriteError::WrongType); + } + + // estimate memory for the new vector + let dim = vector.len(); + let per_vector = + dim * quantization.bytes_per_element() + connectivity * 2 * 8 + element.len() + 80; + let estimated_increase = if is_new { + memory::ENTRY_OVERHEAD + key.len() + VectorSet::BASE_OVERHEAD + per_vector + } else { + per_vector + }; + if !self.enforce_memory_limit(estimated_increase) { + return Err(VectorWriteError::OutOfMemory); + } + + if is_new { + let vs = VectorSet::new(dim, metric, quantization, connectivity, expansion_add) + .map_err(|e| VectorWriteError::IndexError(e.to_string()))?; + let value = Value::Vector(vs); + self.memory.add(key, &value); + self.entries.insert(key.to_owned(), Entry::new(value, None)); + } + + let entry = self + .entries + .get_mut(key) + .expect("just inserted or verified"); + let old_entry_size = memory::entry_size(key, &entry.value); + + let added = match entry.value { + Value::Vector(ref mut vs) => vs + .add(element.clone(), &vector) + .map_err(|e| VectorWriteError::IndexError(e.to_string()))?, + _ => unreachable!(), + }; + entry.touch(); + + let new_entry_size = memory::entry_size(key, &entry.value); + self.memory.adjust(old_entry_size, new_entry_size); + + Ok(VAddResult { + element, + vector, + added, + }) + } + + /// Searches for the k nearest neighbors in a vector set. + #[cfg(feature = "vector")] + pub fn vsim( + &mut self, + key: &str, + query: &[f32], + count: usize, + ef_search: usize, + ) -> Result, WrongType> { + if self.remove_if_expired(key) { + return Ok(Vec::new()); + } + + let entry = match self.entries.get_mut(key) { + Some(e) => e, + None => return Ok(Vec::new()), + }; + + entry.touch(); + + match entry.value { + Value::Vector(ref vs) => vs.search(query, count, ef_search).map_err(|_| WrongType), + _ => Err(WrongType), + } + } + + /// Removes an element from a vector set. Returns `true` if the element + /// existed. Deletes the key if the set becomes empty. + #[cfg(feature = "vector")] + pub fn vrem(&mut self, key: &str, element: &str) -> Result { + if self.remove_if_expired(key) { + return Ok(false); + } + + let entry = match self.entries.get_mut(key) { + Some(e) => e, + None => return Ok(false), + }; + + if !matches!(entry.value, Value::Vector(_)) { + return Err(WrongType); + } + + let old_size = memory::entry_size(key, &entry.value); + + let removed = match entry.value { + Value::Vector(ref mut vs) => vs.remove(element), + _ => unreachable!(), + }; + + if removed { + entry.touch(); + let is_empty = matches!(entry.value, Value::Vector(ref vs) if vs.is_empty()); + let new_size = memory::entry_size(key, &entry.value); + self.memory.adjust(old_size, new_size); + + if is_empty { + self.memory.remove_with_size(new_size); + self.entries.remove(key); + } + } + + Ok(removed) + } + + /// Retrieves the stored vector for an element. + #[cfg(feature = "vector")] + pub fn vget(&mut self, key: &str, element: &str) -> Result>, WrongType> { + if self.remove_if_expired(key) { + return Ok(None); + } + + let entry = match self.entries.get_mut(key) { + Some(e) => e, + None => return Ok(None), + }; + + entry.touch(); + + match entry.value { + Value::Vector(ref vs) => Ok(vs.get(element)), + _ => Err(WrongType), + } + } + + /// Returns the number of elements in a vector set. + #[cfg(feature = "vector")] + pub fn vcard(&mut self, key: &str) -> Result { + if self.remove_if_expired(key) { + return Ok(0); + } + + match self.entries.get(key) { + None => Ok(0), + Some(e) => match e.value { + Value::Vector(ref vs) => Ok(vs.len()), + _ => Err(WrongType), + }, + } + } + + /// Returns the dimensionality of a vector set, or 0 if the key doesn't exist. + #[cfg(feature = "vector")] + pub fn vdim(&mut self, key: &str) -> Result { + if self.remove_if_expired(key) { + return Ok(0); + } + + match self.entries.get(key) { + None => Ok(0), + Some(e) => match e.value { + Value::Vector(ref vs) => Ok(vs.dim()), + _ => Err(WrongType), + }, + } + } + + /// Returns metadata about a vector set. + #[cfg(feature = "vector")] + pub fn vinfo( + &mut self, + key: &str, + ) -> Result, WrongType> { + if self.remove_if_expired(key) { + return Ok(None); + } + + match self.entries.get(key) { + None => Ok(None), + Some(e) => match e.value { + Value::Vector(ref vs) => Ok(Some(vs.info())), + _ => Err(WrongType), + }, + } + } + // -- protobuf operations -- /// Stores a protobuf value. No schema validation here — that's the diff --git a/crates/ember-core/src/lib.rs b/crates/ember-core/src/lib.rs index b09d5df5..ce16145f 100644 --- a/crates/ember-core/src/lib.rs +++ b/crates/ember-core/src/lib.rs @@ -25,5 +25,7 @@ pub use keyspace::{ EvictionPolicy, IncrError, IncrFloatError, Keyspace, KeyspaceStats, RenameError, ShardConfig, TtlResult, WriteError, WrongType, ZAddResult, }; +#[cfg(feature = "vector")] +pub use keyspace::{VAddResult, VectorWriteError}; pub use shard::{ShardPersistenceConfig, ShardRequest, ShardResponse}; pub use types::Value; diff --git a/crates/ember-core/src/memory.rs b/crates/ember-core/src/memory.rs index c2a56b40..ea5d08d6 100644 --- a/crates/ember-core/src/memory.rs +++ b/crates/ember-core/src/memory.rs @@ -160,6 +160,10 @@ pub fn is_large_value(value: &Value) -> bool { Value::SortedSet(ss) => ss.len() > LAZY_FREE_THRESHOLD, Value::Hash(m) => m.len() > LAZY_FREE_THRESHOLD, Value::Set(s) => s.len() > LAZY_FREE_THRESHOLD, + // Vector sets contain usearch Index (C++ object) + hashmaps. + // Large sets should be deferred. + #[cfg(feature = "vector")] + Value::Vector(vs) => vs.len() > LAZY_FREE_THRESHOLD, // Proto values use Bytes (ref-counted, O(1) drop) + a String. // Neither is expensive to drop. #[cfg(feature = "protobuf")] @@ -223,6 +227,8 @@ pub fn value_size(value: &Value) -> usize { let member_bytes: usize = set.iter().map(|m| m.len() + HASHSET_MEMBER_OVERHEAD).sum(); HASHSET_BASE_OVERHEAD + member_bytes } + #[cfg(feature = "vector")] + Value::Vector(vs) => vs.memory_usage(), // type_name: String struct = 24 bytes (ptr+len+cap) on 64-bit. // data: Bytes struct = ~24 bytes (ptr+len+vtable/arc). #[cfg(feature = "protobuf")] diff --git a/crates/ember-core/src/shard.rs b/crates/ember-core/src/shard.rs index 8bdbc5cf..c17dc5e6 100644 --- a/crates/ember-core/src/shard.rs +++ b/crates/ember-core/src/shard.rs @@ -260,6 +260,52 @@ pub enum ShardRequest { slot: u16, count: usize, }, + /// Adds a vector to a vector set. + #[cfg(feature = "vector")] + VAdd { + key: String, + element: String, + vector: Vec, + metric: u8, + quantization: u8, + connectivity: u32, + expansion_add: u32, + }, + /// Searches for nearest neighbors in a vector set. + #[cfg(feature = "vector")] + VSim { + key: String, + query: Vec, + count: usize, + ef_search: usize, + }, + /// Removes an element from a vector set. + #[cfg(feature = "vector")] + VRem { + key: String, + element: String, + }, + /// Gets the stored vector for an element. + #[cfg(feature = "vector")] + VGet { + key: String, + element: String, + }, + /// Returns the number of elements in a vector set. + #[cfg(feature = "vector")] + VCard { + key: String, + }, + /// Returns the dimensionality of a vector set. + #[cfg(feature = "vector")] + VDim { + key: String, + }, + /// Returns metadata about a vector set. + #[cfg(feature = "vector")] + VInfo { + key: String, + }, /// Stores a validated protobuf value. #[cfg(feature = "protobuf")] ProtoSet { @@ -359,6 +405,22 @@ pub enum ShardResponse { StringArray(Vec), /// HMGET result: array of optional values. OptionalArray(Vec>), + /// VADD result: element, vector, and whether it was newly added. + #[cfg(feature = "vector")] + VAddResult { + element: String, + vector: Vec, + added: bool, + }, + /// VSIM result: nearest neighbors with distances. + #[cfg(feature = "vector")] + VSimResult(Vec<(String, f32)>), + /// VGET result: stored vector or None. + #[cfg(feature = "vector")] + VectorData(Option>), + /// VINFO result: vector set metadata. + #[cfg(feature = "vector")] + VectorInfo(Option>), /// PROTO.GET result: (type_name, data, remaining_ttl) or None. #[cfg(feature = "protobuf")] ProtoValue(Option<(String, Bytes, Option)>), @@ -484,6 +546,37 @@ async fn run_shard( } RecoveredValue::Hash(map) => Value::Hash(map), RecoveredValue::Set(set) => Value::Set(set), + #[cfg(feature = "vector")] + RecoveredValue::Vector { + metric, + quantization, + connectivity, + expansion_add, + elements, + } => { + use crate::types::vector::{DistanceMetric, QuantizationType, VectorSet}; + let dim = elements.first().map(|(_, v)| v.len()).unwrap_or(0); + match VectorSet::new( + dim, + DistanceMetric::from_u8(metric), + QuantizationType::from_u8(quantization), + connectivity as usize, + expansion_add as usize, + ) { + Ok(mut vs) => { + for (element, vector) in elements { + if let Err(e) = vs.add(element, &vector) { + warn!("vector recovery: failed to add element: {e}"); + } + } + Value::Vector(vs) + } + Err(e) => { + warn!("vector recovery: failed to create index: {e}"); + continue; + } + } + } #[cfg(feature = "protobuf")] RecoveredValue::Proto { type_name, data } => Value::Proto { type_name, data }, }; @@ -891,6 +984,89 @@ fn dispatch( ShardRequest::GetKeysInSlot { slot, count } => { ShardResponse::StringArray(ks.get_keys_in_slot(*slot, *count)) } + #[cfg(feature = "vector")] + ShardRequest::VAdd { + key, + element, + vector, + metric, + quantization, + connectivity, + expansion_add, + } => { + use crate::types::vector::{DistanceMetric, QuantizationType}; + match ks.vadd( + key, + element.clone(), + vector.clone(), + DistanceMetric::from_u8(*metric), + QuantizationType::from_u8(*quantization), + *connectivity as usize, + *expansion_add as usize, + ) { + Ok(result) => ShardResponse::VAddResult { + element: result.element, + vector: result.vector, + added: result.added, + }, + Err(crate::keyspace::VectorWriteError::WrongType) => ShardResponse::WrongType, + Err(crate::keyspace::VectorWriteError::OutOfMemory) => ShardResponse::OutOfMemory, + Err(crate::keyspace::VectorWriteError::IndexError(e)) => { + ShardResponse::Err(format!("ERR vector index: {e}")) + } + } + } + #[cfg(feature = "vector")] + ShardRequest::VSim { + key, + query, + count, + ef_search, + } => match ks.vsim(key, query, *count, *ef_search) { + Ok(results) => ShardResponse::VSimResult( + results + .into_iter() + .map(|r| (r.element, r.distance)) + .collect(), + ), + Err(_) => ShardResponse::WrongType, + }, + #[cfg(feature = "vector")] + ShardRequest::VRem { key, element } => match ks.vrem(key, element) { + Ok(removed) => ShardResponse::Bool(removed), + Err(_) => ShardResponse::WrongType, + }, + #[cfg(feature = "vector")] + ShardRequest::VGet { key, element } => match ks.vget(key, element) { + Ok(data) => ShardResponse::VectorData(data), + Err(_) => ShardResponse::WrongType, + }, + #[cfg(feature = "vector")] + ShardRequest::VCard { key } => match ks.vcard(key) { + Ok(count) => ShardResponse::Integer(count as i64), + Err(_) => ShardResponse::WrongType, + }, + #[cfg(feature = "vector")] + ShardRequest::VDim { key } => match ks.vdim(key) { + Ok(dim) => ShardResponse::Integer(dim as i64), + Err(_) => ShardResponse::WrongType, + }, + #[cfg(feature = "vector")] + ShardRequest::VInfo { key } => match ks.vinfo(key) { + Ok(Some(info)) => { + let fields = vec![ + ("dim".to_owned(), info.dim.to_string()), + ("count".to_owned(), info.count.to_string()), + ("metric".to_owned(), info.metric.to_string()), + ("quantization".to_owned(), info.quantization.to_string()), + ("connectivity".to_owned(), info.connectivity.to_string()), + ("expansion_add".to_owned(), info.expansion_add.to_string()), + ]; + ShardResponse::VectorInfo(Some(fields)) + } + Ok(None) => ShardResponse::VectorInfo(None), + Err(_) => ShardResponse::WrongType, + }, #[cfg(feature = "protobuf")] ShardRequest::ProtoSet { key, @@ -1204,6 +1380,34 @@ fn to_aof_record(req: &ShardRequest, resp: &ShardResponse) -> Option expire_ms, }) } + // Vector commands + #[cfg(feature = "vector")] + ( + ShardRequest::VAdd { + key, + metric, + quantization, + connectivity, + expansion_add, + .. + }, + ShardResponse::VAddResult { + element, vector, .. + }, + ) => Some(AofRecord::VAdd { + key: key.clone(), + element: element.clone(), + vector: vector.clone(), + metric: *metric, + quantization: *quantization, + connectivity: *connectivity, + expansion_add: *expansion_add, + }), + #[cfg(feature = "vector")] + (ShardRequest::VRem { key, element }, ShardResponse::Bool(true)) => Some(AofRecord::VRem { + key: key.clone(), + element: element.clone(), + }), _ => None, } } @@ -1334,6 +1538,23 @@ fn write_snapshot( } Value::Hash(map) => SnapValue::Hash(map.clone()), Value::Set(set) => SnapValue::Set(set.clone()), + #[cfg(feature = "vector")] + Value::Vector(ref vs) => { + let mut elements = Vec::with_capacity(vs.len()); + for name in vs.elements() { + if let Some(vec) = vs.get(name) { + elements.push((name.to_owned(), vec)); + } + } + SnapValue::Vector { + metric: vs.metric().into(), + quantization: vs.quantization().into(), + connectivity: vs.connectivity() as u32, + expansion_add: vs.expansion_add() as u32, + dim: vs.dim() as u32, + elements, + } + } #[cfg(feature = "protobuf")] Value::Proto { type_name, data } => SnapValue::Proto { type_name: type_name.clone(), diff --git a/crates/ember-core/src/types/mod.rs b/crates/ember-core/src/types/mod.rs index d9da9de2..b20a90d4 100644 --- a/crates/ember-core/src/types/mod.rs +++ b/crates/ember-core/src/types/mod.rs @@ -4,6 +4,8 @@ //! sets, hashes, and sets are supported. pub mod sorted_set; +#[cfg(feature = "vector")] +pub mod vector; use std::collections::{HashMap, HashSet, VecDeque}; @@ -37,6 +39,11 @@ pub enum Value { /// Unordered set of unique string members. Set(HashSet), + /// HNSW-backed vector set for similarity search. Each element is a + /// named string mapped to a dense float vector. + #[cfg(feature = "vector")] + Vector(vector::VectorSet), + /// A protobuf message value. Stores the fully-qualified message type /// name alongside the serialized bytes. Validation happens at the /// server layer before storage. @@ -57,6 +64,8 @@ impl PartialEq for Value { } (Value::Hash(a), Value::Hash(b)) => a == b, (Value::Set(a), Value::Set(b)) => a == b, + #[cfg(feature = "vector")] + (Value::Vector(a), Value::Vector(b)) => a == b, #[cfg(feature = "protobuf")] ( Value::Proto { @@ -81,6 +90,8 @@ pub fn type_name(value: &Value) -> &'static str { Value::SortedSet(_) => "zset", Value::Hash(_) => "hash", Value::Set(_) => "set", + #[cfg(feature = "vector")] + Value::Vector(_) => "vectorset", #[cfg(feature = "protobuf")] Value::Proto { .. } => "proto", } diff --git a/crates/ember-core/src/types/vector.rs b/crates/ember-core/src/types/vector.rs new file mode 100644 index 00000000..0d7ac18d --- /dev/null +++ b/crates/ember-core/src/types/vector.rs @@ -0,0 +1,642 @@ +//! HNSW-backed vector sets for similarity search. +//! +//! Each vector set owns a usearch `Index` with fixed dimensionality and +//! distance metric. Elements are named strings mapped to dense float vectors, +//! analogous to how sorted set members have scores. + +use std::collections::HashMap; +use std::fmt; + +use usearch::{Index, IndexOptions, MetricKind, ScalarKind}; + +/// Distance metric for vector comparison. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum DistanceMetric { + Cosine, + L2, + InnerProduct, +} + +impl DistanceMetric { + fn to_metric_kind(self) -> MetricKind { + match self { + DistanceMetric::Cosine => MetricKind::Cos, + DistanceMetric::L2 => MetricKind::L2sq, + DistanceMetric::InnerProduct => MetricKind::IP, + } + } + + /// Returns the string name used in VINFO output and protocol parsing. + pub fn as_str(self) -> &'static str { + match self { + DistanceMetric::Cosine => "cosine", + DistanceMetric::L2 => "l2", + DistanceMetric::InnerProduct => "ip", + } + } +} + +impl fmt::Display for DistanceMetric { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(self.as_str()) + } +} + +impl From for u8 { + fn from(m: DistanceMetric) -> u8 { + match m { + DistanceMetric::Cosine => 0, + DistanceMetric::L2 => 1, + DistanceMetric::InnerProduct => 2, + } + } +} + +impl DistanceMetric { + /// Converts a wire-format byte to a distance metric. + /// Defaults to `Cosine` for unknown values (forward compatibility). + pub fn from_u8(v: u8) -> Self { + match v { + 1 => DistanceMetric::L2, + 2 => DistanceMetric::InnerProduct, + _ => DistanceMetric::Cosine, + } + } +} + +/// Quantization type for stored vectors. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum QuantizationType { + F32, + F16, + I8, +} + +impl QuantizationType { + fn to_scalar_kind(self) -> ScalarKind { + match self { + QuantizationType::F32 => ScalarKind::F32, + QuantizationType::F16 => ScalarKind::F16, + QuantizationType::I8 => ScalarKind::I8, + } + } + + /// Bytes per element for this quantization level. + pub fn bytes_per_element(self) -> usize { + match self { + QuantizationType::F32 => 4, + QuantizationType::F16 => 2, + QuantizationType::I8 => 1, + } + } + + /// Returns the string name used in VINFO output and protocol parsing. + pub fn as_str(self) -> &'static str { + match self { + QuantizationType::F32 => "f32", + QuantizationType::F16 => "f16", + QuantizationType::I8 => "i8", + } + } +} + +impl fmt::Display for QuantizationType { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(self.as_str()) + } +} + +impl From for u8 { + fn from(q: QuantizationType) -> u8 { + match q { + QuantizationType::F32 => 0, + QuantizationType::F16 => 1, + QuantizationType::I8 => 2, + } + } +} + +impl QuantizationType { + /// Converts a wire-format byte to a quantization type. + /// Defaults to `F32` for unknown values (forward compatibility). + pub fn from_u8(v: u8) -> Self { + match v { + 1 => QuantizationType::F16, + 2 => QuantizationType::I8, + _ => QuantizationType::F32, + } + } +} + +/// Metadata about a vector set, returned by VINFO. +#[derive(Debug, Clone)] +pub struct VectorSetInfo { + pub dim: usize, + pub count: usize, + pub metric: DistanceMetric, + pub quantization: QuantizationType, + pub connectivity: usize, + pub expansion_add: usize, +} + +/// A single search result: element name + distance. +#[derive(Debug, Clone)] +pub struct SearchResult { + pub element: String, + pub distance: f32, +} + +/// Error type for vector operations. +#[derive(Debug, thiserror::Error)] +pub enum VectorError { + #[error("dimension mismatch: index has {expected}, got {got}")] + DimensionMismatch { expected: usize, got: usize }, + + #[error("usearch error: {0}")] + Index(String), +} + +/// A set of named vectors backed by a usearch HNSW index. +/// +/// Each element is a string name mapped to a dense f32 vector. The index +/// is configured on first insert and its parameters (dim, metric, quant, +/// connectivity) are immutable after that. +/// +/// Analogous to Redis sorted sets where members have scores, here elements +/// have vectors. +pub struct VectorSet { + index: Index, + /// element name → usearch key + elements: HashMap, + /// usearch key → element name (for translating search results) + names: HashMap, + /// monotonic key counter for usearch + next_key: u64, + /// vector dimensionality, locked after first insert + dim: usize, + /// distance metric + metric: DistanceMetric, + /// quantization level + quantization: QuantizationType, + /// HNSW connectivity parameter (M) + connectivity: usize, + /// HNSW construction beam width (ef_construction) + expansion_add: usize, +} + +impl VectorSet { + /// Creates a new vector set with the given configuration. + /// + /// `dim` is the fixed dimensionality for all vectors in this set. + /// The usearch index is created eagerly with initial capacity. + pub fn new( + dim: usize, + metric: DistanceMetric, + quantization: QuantizationType, + connectivity: usize, + expansion_add: usize, + ) -> Result { + let options = IndexOptions { + dimensions: dim, + metric: metric.to_metric_kind(), + quantization: quantization.to_scalar_kind(), + connectivity, + expansion_add, + expansion_search: 0, // use default at search time + multi: false, + }; + + let index = Index::new(&options).map_err(|e| VectorError::Index(e.to_string()))?; + index + .reserve(64) + .map_err(|e| VectorError::Index(e.to_string()))?; + + Ok(Self { + index, + elements: HashMap::new(), + names: HashMap::new(), + next_key: 0, + dim, + metric, + quantization, + connectivity, + expansion_add, + }) + } + + /// Adds or replaces a vector for the given element name. + /// + /// Returns `true` if a new element was added, `false` if an existing + /// element was updated. + pub fn add(&mut self, element: String, vector: &[f32]) -> Result { + if vector.len() != self.dim { + return Err(VectorError::DimensionMismatch { + expected: self.dim, + got: vector.len(), + }); + } + + // ensure capacity (double when full, amortized O(1)) + if self.index.size() >= self.index.capacity() { + let new_cap = (self.index.capacity() * 2).max(64); + self.index + .reserve(new_cap) + .map_err(|e| VectorError::Index(e.to_string()))?; + } + + let is_new = if let Some(&existing_key) = self.elements.get(&element) { + // remove old vector, then re-insert with same key + let _ = self.index.remove(existing_key); + self.index + .add(existing_key, vector) + .map_err(|e| VectorError::Index(e.to_string()))?; + false + } else { + let key = self.next_key; + self.next_key += 1; + self.index + .add(key, vector) + .map_err(|e| VectorError::Index(e.to_string()))?; + self.elements.insert(element.clone(), key); + self.names.insert(key, element); + true + }; + + Ok(is_new) + } + + /// Removes an element from the vector set. + /// + /// Returns `true` if the element existed and was removed. + pub fn remove(&mut self, element: &str) -> bool { + if let Some(key) = self.elements.remove(element) { + self.names.remove(&key); + // usearch remove marks the entry as deleted (lazy tombstone). + // the space is reclaimed on subsequent adds. + let _ = self.index.remove(key); + true + } else { + false + } + } + + /// Searches for the k nearest neighbors of the given query vector. + /// + /// `ef_search` controls the search beam width (higher = more accurate, + /// slower). Pass 0 to use usearch's default. + pub fn search( + &self, + query: &[f32], + count: usize, + ef_search: usize, + ) -> Result, VectorError> { + if query.len() != self.dim { + return Err(VectorError::DimensionMismatch { + expected: self.dim, + got: query.len(), + }); + } + + if self.elements.is_empty() { + return Ok(Vec::new()); + } + + // temporarily adjust search expansion if requested + if ef_search > 0 { + self.index.change_expansion_search(ef_search); + } + + let matches = self + .index + .search(query, count) + .map_err(|e| VectorError::Index(e.to_string()))?; + + let mut results = Vec::with_capacity(matches.keys.len()); + for (key, distance) in matches.keys.iter().zip(matches.distances.iter()) { + if let Some(name) = self.names.get(key) { + results.push(SearchResult { + element: name.clone(), + distance: *distance, + }); + } + } + + Ok(results) + } + + /// Retrieves the stored vector for an element. + /// + /// Returns `None` if the element doesn't exist. + pub fn get(&self, element: &str) -> Option> { + let &key = self.elements.get(element)?; + let mut buffer = vec![0.0f32; self.dim]; + match self.index.get(key, &mut buffer) { + Ok(found) if found > 0 => Some(buffer), + _ => None, + } + } + + /// Returns the number of elements in the vector set. + pub fn len(&self) -> usize { + self.elements.len() + } + + /// Returns `true` if the vector set has no elements. + pub fn is_empty(&self) -> bool { + self.elements.is_empty() + } + + /// Returns the dimensionality of vectors in this set. + pub fn dim(&self) -> usize { + self.dim + } + + /// Returns the distance metric. + pub fn metric(&self) -> DistanceMetric { + self.metric + } + + /// Returns the quantization type. + pub fn quantization(&self) -> QuantizationType { + self.quantization + } + + /// Returns metadata about this vector set. + pub fn info(&self) -> VectorSetInfo { + VectorSetInfo { + dim: self.dim, + count: self.elements.len(), + metric: self.metric, + quantization: self.quantization, + connectivity: self.connectivity, + expansion_add: self.expansion_add, + } + } + + /// Returns an iterator over all element names. + /// + /// Used for snapshot serialization — the caller retrieves each vector + /// via `get()`. + pub fn elements(&self) -> impl Iterator { + self.elements.keys().map(String::as_str) + } + + /// Returns the HNSW connectivity parameter. + pub fn connectivity(&self) -> usize { + self.connectivity + } + + /// Returns the HNSW construction beam width. + pub fn expansion_add(&self) -> usize { + self.expansion_add + } + + /// Estimates memory usage in bytes. + /// + /// Accounts for: usearch index storage (vectors + HNSW graph), + /// element↔key hashmaps, and string names. + pub fn memory_usage(&self) -> usize { + let count = self.elements.len(); + + // usearch internal: vector storage + HNSW graph edges + let vector_bytes = count * self.dim * self.quantization.bytes_per_element(); + let graph_bytes = count * self.connectivity * 2 * 8; // each edge is a u64 key + + // rust-side hashmaps: elements + names + let name_bytes: usize = self + .elements + .keys() + .map(|name| name.len() + 80) // String + HashMap entry overhead for both maps + .sum(); + + Self::BASE_OVERHEAD + vector_bytes + graph_bytes + name_bytes + } + + /// Base overhead of an empty VectorSet (usearch index shell + two HashMaps). + pub const BASE_OVERHEAD: usize = 128; +} + +impl fmt::Debug for VectorSet { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("VectorSet") + .field("dim", &self.dim) + .field("count", &self.elements.len()) + .field("metric", &self.metric) + .field("quantization", &self.quantization) + .finish() + } +} + +impl Clone for VectorSet { + fn clone(&self) -> Self { + // rebuild the index from scratch — usearch Index doesn't implement Clone + let mut new = Self::new( + self.dim, + self.metric, + self.quantization, + self.connectivity, + self.expansion_add, + ) + .expect("clone: failed to create index with same config"); + + for (name, &key) in &self.elements { + let mut buffer = vec![0.0f32; self.dim]; + if self.index.get(key, &mut buffer).is_ok() { + let _ = new.add(name.clone(), &buffer); + } + } + + new + } +} + +impl PartialEq for VectorSet { + fn eq(&self, other: &Self) -> bool { + if self.dim != other.dim + || self.metric != other.metric + || self.quantization != other.quantization + || self.elements.len() != other.elements.len() + { + return false; + } + + // compare all element names and their vectors + for (name, &key) in &self.elements { + match other.elements.get(name) { + Some(&other_key) => { + let mut buf_a = vec![0.0f32; self.dim]; + let mut buf_b = vec![0.0f32; self.dim]; + let ok_a = self.index.get(key, &mut buf_a).is_ok(); + let ok_b = other.index.get(other_key, &mut buf_b).is_ok(); + if ok_a != ok_b || buf_a != buf_b { + return false; + } + } + None => return false, + } + } + + true + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn make_set(dim: usize) -> VectorSet { + VectorSet::new(dim, DistanceMetric::Cosine, QuantizationType::F32, 16, 64).unwrap() + } + + #[test] + fn add_and_get() { + let mut vs = make_set(3); + let added = vs.add("a".into(), &[1.0, 0.0, 0.0]).unwrap(); + assert!(added); + assert_eq!(vs.len(), 1); + + let vec = vs.get("a").unwrap(); + assert_eq!(vec, vec![1.0, 0.0, 0.0]); + } + + #[test] + fn add_update_existing() { + let mut vs = make_set(3); + vs.add("a".into(), &[1.0, 0.0, 0.0]).unwrap(); + + let added = vs.add("a".into(), &[0.0, 1.0, 0.0]).unwrap(); + assert!(!added); // update, not new + assert_eq!(vs.len(), 1); + + let vec = vs.get("a").unwrap(); + assert_eq!(vec, vec![0.0, 1.0, 0.0]); + } + + #[test] + fn dimension_mismatch() { + let mut vs = make_set(3); + let err = vs.add("a".into(), &[1.0, 0.0]).unwrap_err(); + assert!(matches!( + err, + VectorError::DimensionMismatch { + expected: 3, + got: 2 + } + )); + } + + #[test] + fn remove_element() { + let mut vs = make_set(3); + vs.add("a".into(), &[1.0, 0.0, 0.0]).unwrap(); + + assert!(vs.remove("a")); + assert_eq!(vs.len(), 0); + assert!(vs.get("a").is_none()); + assert!(!vs.remove("a")); // already gone + } + + #[test] + fn search_basic() { + let mut vs = make_set(3); + vs.add("x-axis".into(), &[1.0, 0.0, 0.0]).unwrap(); + vs.add("y-axis".into(), &[0.0, 1.0, 0.0]).unwrap(); + vs.add("z-axis".into(), &[0.0, 0.0, 1.0]).unwrap(); + + // searching near x-axis should return x-axis first + let results = vs.search(&[0.9, 0.1, 0.0], 2, 0).unwrap(); + assert_eq!(results.len(), 2); + assert_eq!(results[0].element, "x-axis"); + } + + #[test] + fn search_empty_set() { + let vs = make_set(3); + let results = vs.search(&[1.0, 0.0, 0.0], 5, 0).unwrap(); + assert!(results.is_empty()); + } + + #[test] + fn search_dimension_mismatch() { + let vs = make_set(3); + let err = vs.search(&[1.0, 0.0], 5, 0).unwrap_err(); + assert!(matches!(err, VectorError::DimensionMismatch { .. })); + } + + #[test] + fn get_nonexistent() { + let vs = make_set(3); + assert!(vs.get("nope").is_none()); + } + + #[test] + fn info() { + let vs = make_set(4); + let info = vs.info(); + assert_eq!(info.dim, 4); + assert_eq!(info.count, 0); + assert_eq!(info.metric, DistanceMetric::Cosine); + assert_eq!(info.quantization, QuantizationType::F32); + assert_eq!(info.connectivity, 16); + assert_eq!(info.expansion_add, 64); + } + + #[test] + fn memory_usage_grows() { + let mut vs = make_set(128); + let base = vs.memory_usage(); + + vs.add("a".into(), &vec![0.0; 128]).unwrap(); + let with_one = vs.memory_usage(); + assert!(with_one > base); + + vs.add("b".into(), &vec![0.0; 128]).unwrap(); + assert!(vs.memory_usage() > with_one); + } + + #[test] + fn clone_preserves_data() { + let mut vs = make_set(3); + vs.add("a".into(), &[1.0, 2.0, 3.0]).unwrap(); + vs.add("b".into(), &[4.0, 5.0, 6.0]).unwrap(); + + let cloned = vs.clone(); + assert_eq!(cloned.len(), 2); + assert_eq!(cloned.get("a").unwrap(), vec![1.0, 2.0, 3.0]); + assert_eq!(cloned.get("b").unwrap(), vec![4.0, 5.0, 6.0]); + } + + #[test] + fn partial_eq() { + let mut a = make_set(3); + a.add("x".into(), &[1.0, 0.0, 0.0]).unwrap(); + + let mut b = make_set(3); + b.add("x".into(), &[1.0, 0.0, 0.0]).unwrap(); + + assert_eq!(a, b); + } + + #[test] + fn l2_metric() { + let mut vs = VectorSet::new(2, DistanceMetric::L2, QuantizationType::F32, 16, 64).unwrap(); + vs.add("origin".into(), &[0.0, 0.0]).unwrap(); + vs.add("far".into(), &[10.0, 10.0]).unwrap(); + + let results = vs.search(&[0.1, 0.1], 1, 0).unwrap(); + assert_eq!(results[0].element, "origin"); + } + + #[test] + fn different_quantization() { + // F16 should work, though values may lose precision + let mut vs = + VectorSet::new(3, DistanceMetric::Cosine, QuantizationType::F16, 16, 64).unwrap(); + vs.add("a".into(), &[1.0, 0.0, 0.0]).unwrap(); + assert_eq!(vs.len(), 1); + + // we can still get the vector back (with possible precision loss) + let vec = vs.get("a").unwrap(); + assert!((vec[0] - 1.0).abs() < 0.01); + } +} diff --git a/crates/ember-persistence/Cargo.toml b/crates/ember-persistence/Cargo.toml index 0ca5c0d6..9e32dc7f 100644 --- a/crates/ember-persistence/Cargo.toml +++ b/crates/ember-persistence/Cargo.toml @@ -12,6 +12,7 @@ readme = "README.md" [features] encryption = ["aes-gcm", "rand"] protobuf = [] +vector = [] [dependencies] thiserror = { workspace = true } diff --git a/crates/ember-persistence/src/aof.rs b/crates/ember-persistence/src/aof.rs index c7317a7b..232e886c 100644 --- a/crates/ember-persistence/src/aof.rs +++ b/crates/ember-persistence/src/aof.rs @@ -86,6 +86,12 @@ const TAG_PERSIST: u8 = 10; const TAG_PEXPIRE: u8 = 11; const TAG_RENAME: u8 = 22; +// vector +#[cfg(feature = "vector")] +const TAG_VADD: u8 = 25; +#[cfg(feature = "vector")] +const TAG_VREM: u8 = 26; + // protobuf #[cfg(feature = "protobuf")] const TAG_PROTO_SET: u8 = 23; @@ -153,6 +159,23 @@ pub enum AofRecord { Append { key: String, value: Bytes }, /// RENAME key newkey. Rename { key: String, newkey: String }, + /// VADD key element vector [metric quant connectivity expansion_add]. + /// Stores the full index config so recovery can recreate the set. + #[cfg(feature = "vector")] + VAdd { + key: String, + element: String, + vector: Vec, + /// 0 = cosine, 1 = l2, 2 = inner product + metric: u8, + /// 0 = f32, 1 = f16, 2 = i8 + quantization: u8, + connectivity: u32, + expansion_add: u32, + }, + /// VREM key element. + #[cfg(feature = "vector")] + VRem { key: String, element: String }, /// PROTO.SET key type_name data [expire_ms]. #[cfg(feature = "protobuf")] ProtoSet { @@ -307,6 +330,34 @@ impl AofRecord { format::write_bytes(&mut buf, key.as_bytes())?; format::write_bytes(&mut buf, newkey.as_bytes())?; } + #[cfg(feature = "vector")] + AofRecord::VAdd { + key, + element, + vector, + metric, + quantization, + connectivity, + expansion_add, + } => { + format::write_u8(&mut buf, TAG_VADD)?; + format::write_bytes(&mut buf, key.as_bytes())?; + format::write_bytes(&mut buf, element.as_bytes())?; + format::write_u32(&mut buf, vector.len() as u32)?; + for &v in vector { + format::write_f32(&mut buf, v)?; + } + format::write_u8(&mut buf, *metric)?; + format::write_u8(&mut buf, *quantization)?; + format::write_u32(&mut buf, *connectivity)?; + format::write_u32(&mut buf, *expansion_add)?; + } + #[cfg(feature = "vector")] + AofRecord::VRem { key, element } => { + format::write_u8(&mut buf, TAG_VREM)?; + format::write_bytes(&mut buf, key.as_bytes())?; + format::write_bytes(&mut buf, element.as_bytes())?; + } #[cfg(feature = "protobuf")] AofRecord::ProtoSet { key, @@ -460,6 +511,35 @@ impl AofRecord { let newkey = read_string(&mut cursor, "newkey")?; Ok(AofRecord::Rename { key, newkey }) } + #[cfg(feature = "vector")] + TAG_VADD => { + let key = read_string(&mut cursor, "key")?; + let element = read_string(&mut cursor, "element")?; + let dim = format::read_u32(&mut cursor)?; + let mut vector = Vec::with_capacity(format::capped_capacity(dim)); + for _ in 0..dim { + vector.push(format::read_f32(&mut cursor)?); + } + let metric = format::read_u8(&mut cursor)?; + let quantization = format::read_u8(&mut cursor)?; + let connectivity = format::read_u32(&mut cursor)?; + let expansion_add = format::read_u32(&mut cursor)?; + Ok(AofRecord::VAdd { + key, + element, + vector, + metric, + quantization, + connectivity, + expansion_add, + }) + } + #[cfg(feature = "vector")] + TAG_VREM => { + let key = read_string(&mut cursor, "key")?; + let element = read_string(&mut cursor, "element")?; + Ok(AofRecord::VRem { key, element }) + } #[cfg(feature = "protobuf")] TAG_PROTO_SET => { let key = read_string(&mut cursor, "key")?; @@ -903,6 +983,34 @@ impl AofReader { let newkey = format::read_bytes(&mut self.reader)?; format::write_bytes(&mut payload, &newkey)?; } + #[cfg(feature = "vector")] + TAG_VADD => { + let key = format::read_bytes(&mut self.reader)?; + format::write_bytes(&mut payload, &key)?; + let element = format::read_bytes(&mut self.reader)?; + format::write_bytes(&mut payload, &element)?; + let dim = format::read_u32(&mut self.reader)?; + format::write_u32(&mut payload, dim)?; + for _ in 0..dim { + let v = format::read_f32(&mut self.reader)?; + format::write_f32(&mut payload, v)?; + } + let metric = format::read_u8(&mut self.reader)?; + format::write_u8(&mut payload, metric)?; + let quantization = format::read_u8(&mut self.reader)?; + format::write_u8(&mut payload, quantization)?; + let connectivity = format::read_u32(&mut self.reader)?; + format::write_u32(&mut payload, connectivity)?; + let expansion_add = format::read_u32(&mut self.reader)?; + format::write_u32(&mut payload, expansion_add)?; + } + #[cfg(feature = "vector")] + TAG_VREM => { + let key = format::read_bytes(&mut self.reader)?; + format::write_bytes(&mut payload, &key)?; + let element = format::read_bytes(&mut self.reader)?; + format::write_bytes(&mut payload, &element)?; + } #[cfg(feature = "protobuf")] TAG_PROTO_SET => { let key = format::read_bytes(&mut self.reader)?; diff --git a/crates/ember-persistence/src/format.rs b/crates/ember-persistence/src/format.rs index 1b6a75cc..6d3fc115 100644 --- a/crates/ember-persistence/src/format.rs +++ b/crates/ember-persistence/src/format.rs @@ -84,6 +84,11 @@ pub fn write_i64(w: &mut impl Write, val: i64) -> io::Result<()> { w.write_all(&val.to_le_bytes()) } +/// Writes an `f32` in little-endian. +pub fn write_f32(w: &mut impl Write, val: f32) -> io::Result<()> { + w.write_all(&val.to_le_bytes()) +} + /// Writes an `f64` in little-endian. pub fn write_f64(w: &mut impl Write, val: f64) -> io::Result<()> { w.write_all(&val.to_le_bytes()) @@ -127,6 +132,13 @@ pub fn read_i64(r: &mut impl Read) -> Result { Ok(i64::from_le_bytes(buf)) } +/// Reads an `f32` in little-endian. +pub fn read_f32(r: &mut impl Read) -> Result { + let mut buf = [0u8; 4]; + read_exact(r, &mut buf)?; + Ok(f32::from_le_bytes(buf)) +} + /// Reads an `f64` in little-endian. pub fn read_f64(r: &mut impl Read) -> Result { let mut buf = [0u8; 8]; diff --git a/crates/ember-persistence/src/recovery.rs b/crates/ember-persistence/src/recovery.rs index 90ddb831..ca845e04 100644 --- a/crates/ember-persistence/src/recovery.rs +++ b/crates/ember-persistence/src/recovery.rs @@ -37,6 +37,15 @@ pub enum RecoveredValue { Hash(HashMap), /// Unordered set of unique string members. Set(HashSet), + /// A vector set: index config + accumulated (element, vector) pairs. + #[cfg(feature = "vector")] + Vector { + metric: u8, + quantization: u8, + connectivity: u32, + expansion_add: u32, + elements: Vec<(String, Vec)>, + }, /// A protobuf message: type name + serialized bytes. #[cfg(feature = "protobuf")] Proto { @@ -53,6 +62,21 @@ impl From for RecoveredValue { SnapValue::SortedSet(members) => RecoveredValue::SortedSet(members), SnapValue::Hash(map) => RecoveredValue::Hash(map), SnapValue::Set(set) => RecoveredValue::Set(set), + #[cfg(feature = "vector")] + SnapValue::Vector { + metric, + quantization, + connectivity, + expansion_add, + elements, + .. + } => RecoveredValue::Vector { + metric, + quantization, + connectivity, + expansion_add, + elements, + }, #[cfg(feature = "protobuf")] SnapValue::Proto { type_name, data } => RecoveredValue::Proto { type_name, data }, } @@ -441,6 +465,54 @@ fn replay_aof( } } } + #[cfg(feature = "vector")] + AofRecord::VAdd { + key, + element, + vector, + metric, + quantization, + connectivity, + expansion_add, + } => { + let entry = map.entry(key).or_insert_with(|| { + ( + RecoveredValue::Vector { + metric, + quantization, + connectivity, + expansion_add, + elements: Vec::new(), + }, + -1, // no expiry for vector sets + ) + }); + if let RecoveredValue::Vector { + ref mut elements, .. + } = entry.0 + { + // replace existing element or add new + if let Some(pos) = elements.iter().position(|(e, _)| *e == element) { + elements[pos].1 = vector; + } else { + elements.push((element, vector)); + } + } + } + #[cfg(feature = "vector")] + AofRecord::VRem { key, element } => { + if let Some(entry) = map.get_mut(&key) { + if let RecoveredValue::Vector { + ref mut elements, .. + } = entry.0 + { + elements.retain(|(e, _)| *e != element); + if elements.is_empty() { + map.remove(&key); + } + } + } + } #[cfg(feature = "protobuf")] AofRecord::ProtoSet { key, diff --git a/crates/ember-persistence/src/snapshot.rs b/crates/ember-persistence/src/snapshot.rs index 20291dd2..41a923c9 100644 --- a/crates/ember-persistence/src/snapshot.rs +++ b/crates/ember-persistence/src/snapshot.rs @@ -36,6 +36,8 @@ const TYPE_LIST: u8 = 1; const TYPE_SORTED_SET: u8 = 2; const TYPE_HASH: u8 = 3; const TYPE_SET: u8 = 4; +#[cfg(feature = "vector")] +const TYPE_VECTOR: u8 = 6; #[cfg(feature = "protobuf")] const TYPE_PROTO: u8 = 5; @@ -101,6 +103,32 @@ fn parse_snap_value(r: &mut impl io::Read) -> Result { } Ok(SnapValue::Set(set)) } + #[cfg(feature = "vector")] + TYPE_VECTOR => { + let metric = format::read_u8(r)?; + let quantization = format::read_u8(r)?; + let connectivity = format::read_u32(r)?; + let expansion_add = format::read_u32(r)?; + let dim = format::read_u32(r)?; + let count = format::read_u32(r)?; + let mut elements = Vec::with_capacity(format::capped_capacity(count)); + for _ in 0..count { + let name = read_snap_string(r, "vector element name")?; + let mut vector = Vec::with_capacity(format::capped_capacity(dim)); + for _ in 0..dim { + vector.push(format::read_f32(r)?); + } + elements.push((name, vector)); + } + Ok(SnapValue::Vector { + metric, + quantization, + connectivity, + expansion_add, + dim, + elements, + }) + } #[cfg(feature = "protobuf")] TYPE_PROTO => { let type_name = read_snap_string(r, "proto type_name")?; @@ -127,6 +155,16 @@ pub enum SnapValue { Hash(HashMap), /// An unordered set of unique string members. Set(HashSet), + /// A vector set: index config + all (element, vector) pairs. + #[cfg(feature = "vector")] + Vector { + metric: u8, + quantization: u8, + connectivity: u32, + expansion_add: u32, + dim: u32, + elements: Vec<(String, Vec)>, + }, /// A protobuf message: type name + serialized bytes. #[cfg(feature = "protobuf")] Proto { type_name: String, data: Bytes }, @@ -270,6 +308,29 @@ impl SnapshotWriter { format::write_bytes(&mut buf, member.as_bytes())?; } } + #[cfg(feature = "vector")] + SnapValue::Vector { + metric, + quantization, + connectivity, + expansion_add, + dim, + elements, + } => { + format::write_u8(&mut buf, TYPE_VECTOR)?; + format::write_u8(&mut buf, *metric)?; + format::write_u8(&mut buf, *quantization)?; + format::write_u32(&mut buf, *connectivity)?; + format::write_u32(&mut buf, *expansion_add)?; + format::write_u32(&mut buf, *dim)?; + format::write_u32(&mut buf, elements.len() as u32)?; + for (name, vector) in elements { + format::write_bytes(&mut buf, name.as_bytes())?; + for &v in vector { + format::write_f32(&mut buf, v)?; + } + } + } #[cfg(feature = "protobuf")] SnapValue::Proto { type_name, data } => { format::write_u8(&mut buf, TYPE_PROTO)?; @@ -513,6 +574,47 @@ impl SnapshotReader { } SnapValue::Set(set) } + #[cfg(feature = "vector")] + TYPE_VECTOR => { + let metric = format::read_u8(&mut self.reader)?; + format::write_u8(&mut buf, metric)?; + let quantization = format::read_u8(&mut self.reader)?; + format::write_u8(&mut buf, quantization)?; + let connectivity = format::read_u32(&mut self.reader)?; + format::write_u32(&mut buf, connectivity)?; + let expansion_add = format::read_u32(&mut self.reader)?; + format::write_u32(&mut buf, expansion_add)?; + let dim = format::read_u32(&mut self.reader)?; + format::write_u32(&mut buf, dim)?; + let count = format::read_u32(&mut self.reader)?; + format::write_u32(&mut buf, count)?; + let mut elements = Vec::with_capacity(format::capped_capacity(count)); + for _ in 0..count { + let name_bytes = format::read_bytes(&mut self.reader)?; + format::write_bytes(&mut buf, &name_bytes)?; + let name = String::from_utf8(name_bytes).map_err(|_| { + FormatError::Io(io::Error::new( + io::ErrorKind::InvalidData, + "vector element name is not valid utf-8", + )) + })?; + let mut vector = Vec::with_capacity(format::capped_capacity(dim)); + for _ in 0..dim { + let v = format::read_f32(&mut self.reader)?; + format::write_f32(&mut buf, v)?; + vector.push(v); + } + elements.push((name, vector)); + } + SnapValue::Vector { + metric, + quantization, + connectivity, + expansion_add, + dim, + elements, + } + } #[cfg(feature = "protobuf")] TYPE_PROTO => { let type_name_bytes = format::read_bytes(&mut self.reader)?; diff --git a/crates/ember-protocol/src/command.rs b/crates/ember-protocol/src/command.rs index 925ada21..4b62b8f8 100644 --- a/crates/ember-protocol/src/command.rs +++ b/crates/ember-protocol/src/command.rs @@ -321,6 +321,48 @@ pub enum Command { /// PUBSUB NUMPAT. Returns the number of active pattern subscriptions. PubSubNumPat, + // --- vector commands --- + /// VADD key element f32 [f32 ...] [METRIC COSINE|L2|IP] [QUANT F32|F16|I8] + /// [M n] [EF n]. Adds a vector to a vector set. + VAdd { + key: String, + element: String, + vector: Vec, + /// 0 = cosine (default), 1 = l2, 2 = inner product + metric: u8, + /// 0 = f32 (default), 1 = f16, 2 = i8 + quantization: u8, + /// HNSW connectivity parameter (default 16) + connectivity: u32, + /// HNSW construction beam width (default 64) + expansion_add: u32, + }, + + /// VSIM key f32 [f32 ...] COUNT k [EF n] [WITHSCORES]. + /// Searches for k nearest neighbors. + VSim { + key: String, + query: Vec, + count: usize, + ef_search: usize, + with_scores: bool, + }, + + /// VREM key element. Removes a vector from a vector set. + VRem { key: String, element: String }, + + /// VGET key element. Retrieves the stored vector for an element. + VGet { key: String, element: String }, + + /// VCARD key. Returns the number of elements in a vector set. + VCard { key: String }, + + /// VDIM key. Returns the dimensionality of a vector set. + VDim { key: String }, + + /// VINFO key. Returns metadata about a vector set. + VInfo { key: String }, + // --- protobuf commands --- /// PROTO.REGISTER `name` `descriptor_bytes`. Registers a protobuf schema /// (pre-compiled FileDescriptorSet) under the given name. @@ -492,6 +534,13 @@ impl Command { Command::PubSubChannels { .. } => "pubsub", Command::PubSubNumSub { .. } => "pubsub", Command::PubSubNumPat => "pubsub", + Command::VAdd { .. } => "vadd", + Command::VSim { .. } => "vsim", + Command::VRem { .. } => "vrem", + Command::VGet { .. } => "vget", + Command::VCard { .. } => "vcard", + Command::VDim { .. } => "vdim", + Command::VInfo { .. } => "vinfo", Command::ProtoRegister { .. } => "proto.register", Command::ProtoSet { .. } => "proto.set", Command::ProtoGet { .. } => "proto.get", @@ -598,6 +647,13 @@ impl Command { "PUNSUBSCRIBE" => parse_punsubscribe(&frames[1..]), "PUBLISH" => parse_publish(&frames[1..]), "PUBSUB" => parse_pubsub(&frames[1..]), + "VADD" => parse_vadd(&frames[1..]), + "VSIM" => parse_vsim(&frames[1..]), + "VREM" => parse_vrem(&frames[1..]), + "VGET" => parse_vget(&frames[1..]), + "VCARD" => parse_vcard(&frames[1..]), + "VDIM" => parse_vdim(&frames[1..]), + "VINFO" => parse_vinfo(&frames[1..]), "PROTO.REGISTER" => parse_proto_register(&frames[1..]), "PROTO.SET" => parse_proto_set(&frames[1..]), "PROTO.GET" => parse_proto_get(&frames[1..]), @@ -1757,6 +1813,252 @@ fn parse_pubsub(args: &[Frame]) -> Result { } } +// --- vector command parsers --- + +/// VADD key element f32 [f32 ...] [METRIC COSINE|L2|IP] [QUANT F32|F16|I8] [M n] [EF n] +fn parse_vadd(args: &[Frame]) -> Result { + // minimum: key + element + at least one float + if args.len() < 3 { + return Err(ProtocolError::WrongArity("VADD".into())); + } + + let key = extract_string(&args[0])?; + let element = extract_string(&args[1])?; + + // parse vector values until we hit a non-numeric argument or end + let mut idx = 2; + let mut vector = Vec::new(); + while idx < args.len() { + let s = extract_string(&args[idx])?; + if let Ok(v) = s.parse::() { + vector.push(v); + idx += 1; + } else { + break; + } + } + + if vector.is_empty() { + return Err(ProtocolError::InvalidCommandFrame( + "VADD: at least one vector dimension required".into(), + )); + } + + // parse optional flags + let mut metric: u8 = 0; // cosine default + let mut quantization: u8 = 0; // f32 default + let mut connectivity: u32 = 16; + let mut expansion_add: u32 = 64; + + while idx < args.len() { + let flag = extract_string(&args[idx])?.to_ascii_uppercase(); + match flag.as_str() { + "METRIC" => { + idx += 1; + if idx >= args.len() { + return Err(ProtocolError::InvalidCommandFrame( + "VADD: METRIC requires a value".into(), + )); + } + let val = extract_string(&args[idx])?.to_ascii_uppercase(); + metric = match val.as_str() { + "COSINE" => 0, + "L2" => 1, + "IP" => 2, + _ => { + return Err(ProtocolError::InvalidCommandFrame(format!( + "VADD: unknown metric '{val}'" + ))) + } + }; + idx += 1; + } + "QUANT" => { + idx += 1; + if idx >= args.len() { + return Err(ProtocolError::InvalidCommandFrame( + "VADD: QUANT requires a value".into(), + )); + } + let val = extract_string(&args[idx])?.to_ascii_uppercase(); + quantization = match val.as_str() { + "F32" => 0, + "F16" => 1, + "I8" | "Q8" => 2, + _ => { + return Err(ProtocolError::InvalidCommandFrame(format!( + "VADD: unknown quantization '{val}'" + ))) + } + }; + idx += 1; + } + "M" => { + idx += 1; + if idx >= args.len() { + return Err(ProtocolError::InvalidCommandFrame( + "VADD: M requires a value".into(), + )); + } + connectivity = parse_u64(&args[idx], "VADD")? as u32; + idx += 1; + } + "EF" => { + idx += 1; + if idx >= args.len() { + return Err(ProtocolError::InvalidCommandFrame( + "VADD: EF requires a value".into(), + )); + } + expansion_add = parse_u64(&args[idx], "VADD")? as u32; + idx += 1; + } + _ => { + return Err(ProtocolError::InvalidCommandFrame(format!( + "VADD: unexpected argument '{flag}'" + ))); + } + } + } + + Ok(Command::VAdd { + key, + element, + vector, + metric, + quantization, + connectivity, + expansion_add, + }) +} + +/// VSIM key f32 [f32 ...] COUNT k [EF n] [WITHSCORES] +fn parse_vsim(args: &[Frame]) -> Result { + // minimum: key + at least one float + COUNT + k + if args.len() < 4 { + return Err(ProtocolError::WrongArity("VSIM".into())); + } + + let key = extract_string(&args[0])?; + + // parse query vector until we hit a non-numeric argument + let mut idx = 1; + let mut query = Vec::new(); + while idx < args.len() { + let s = extract_string(&args[idx])?; + if let Ok(v) = s.parse::() { + query.push(v); + idx += 1; + } else { + break; + } + } + + if query.is_empty() { + return Err(ProtocolError::InvalidCommandFrame( + "VSIM: at least one query dimension required".into(), + )); + } + + // COUNT k is required + let mut count: Option = None; + let mut ef_search: usize = 0; + let mut with_scores = false; + + while idx < args.len() { + let flag = extract_string(&args[idx])?.to_ascii_uppercase(); + match flag.as_str() { + "COUNT" => { + idx += 1; + if idx >= args.len() { + return Err(ProtocolError::InvalidCommandFrame( + "VSIM: COUNT requires a value".into(), + )); + } + count = Some(parse_u64(&args[idx], "VSIM")? as usize); + idx += 1; + } + "EF" => { + idx += 1; + if idx >= args.len() { + return Err(ProtocolError::InvalidCommandFrame( + "VSIM: EF requires a value".into(), + )); + } + ef_search = parse_u64(&args[idx], "VSIM")? as usize; + idx += 1; + } + "WITHSCORES" => { + with_scores = true; + idx += 1; + } + _ => { + return Err(ProtocolError::InvalidCommandFrame(format!( + "VSIM: unexpected argument '{flag}'" + ))); + } + } + } + + let count = count + .ok_or_else(|| ProtocolError::InvalidCommandFrame("VSIM: COUNT is required".into()))?; + + Ok(Command::VSim { + key, + query, + count, + ef_search, + with_scores, + }) +} + +/// VREM key element +fn parse_vrem(args: &[Frame]) -> Result { + if args.len() != 2 { + return Err(ProtocolError::WrongArity("VREM".into())); + } + let key = extract_string(&args[0])?; + let element = extract_string(&args[1])?; + Ok(Command::VRem { key, element }) +} + +/// VGET key element +fn parse_vget(args: &[Frame]) -> Result { + if args.len() != 2 { + return Err(ProtocolError::WrongArity("VGET".into())); + } + let key = extract_string(&args[0])?; + let element = extract_string(&args[1])?; + Ok(Command::VGet { key, element }) +} + +/// VCARD key +fn parse_vcard(args: &[Frame]) -> Result { + if args.len() != 1 { + return Err(ProtocolError::WrongArity("VCARD".into())); + } + let key = extract_string(&args[0])?; + Ok(Command::VCard { key }) +} + +/// VDIM key +fn parse_vdim(args: &[Frame]) -> Result { + if args.len() != 1 { + return Err(ProtocolError::WrongArity("VDIM".into())); + } + let key = extract_string(&args[0])?; + Ok(Command::VDim { key }) +} + +/// VINFO key +fn parse_vinfo(args: &[Frame]) -> Result { + if args.len() != 1 { + return Err(ProtocolError::WrongArity("VINFO".into())); + } + let key = extract_string(&args[0])?; + Ok(Command::VInfo { key }) +} + // --- proto command parsers --- fn parse_proto_register(args: &[Frame]) -> Result { diff --git a/crates/ember-server/Cargo.toml b/crates/ember-server/Cargo.toml index a1472c3c..59ead69f 100644 --- a/crates/ember-server/Cargo.toml +++ b/crates/ember-server/Cargo.toml @@ -14,6 +14,7 @@ default = ["jemalloc"] jemalloc = ["tikv-jemallocator"] encryption = ["emberkv-core/encryption", "ember-persistence/encryption"] protobuf = ["emberkv-core/protobuf", "ember-persistence/protobuf"] +vector = ["emberkv-core/vector", "ember-persistence/vector"] [dependencies] bytes = { workspace = true } diff --git a/crates/ember-server/src/connection.rs b/crates/ember-server/src/connection.rs index f06819d0..4e015c15 100644 --- a/crates/ember-server/src/connection.rs +++ b/crates/ember-server/src/connection.rs @@ -573,7 +573,14 @@ async fn cluster_slot_check(ctx: &ServerContext, cmd: &Command) -> Option | Command::ProtoType { ref key } | Command::ProtoGetField { ref key, .. } | Command::ProtoSetField { ref key, .. } - | Command::ProtoDelField { ref key, .. } => cluster.check_slot(key.as_bytes()).await, + | Command::ProtoDelField { ref key, .. } + | Command::VAdd { ref key, .. } + | Command::VSim { ref key, .. } + | Command::VRem { ref key, .. } + | Command::VGet { ref key, .. } + | Command::VCard { ref key } + | Command::VDim { ref key } + | Command::VInfo { ref key } => cluster.check_slot(key.as_bytes()).await, // multi-key commands — crossslot validation + slot ownership Command::Del { ref keys } @@ -1644,7 +1651,155 @@ async fn execute( Frame::Error("ERR subscribe commands should not reach execute".into()) } - // -- protobuf commands -- + // --- vector commands --- + #[cfg(feature = "vector")] + Command::VAdd { + key, + element, + vector, + metric, + quantization, + connectivity, + expansion_add, + } => { + let req = ShardRequest::VAdd { + key: key.clone(), + element, + vector, + metric, + quantization, + connectivity, + expansion_add, + }; + match engine.route(&key, req).await { + Ok(ShardResponse::VAddResult { added, .. }) => { + Frame::Integer(if added { 1 } else { 0 }) + } + Ok(ShardResponse::WrongType) => wrongtype_error(), + Ok(ShardResponse::OutOfMemory) => oom_error(), + Ok(ShardResponse::Err(msg)) => Frame::Error(format!("ERR {msg}")), + Ok(other) => Frame::Error(format!("ERR unexpected shard response: {other:?}")), + Err(e) => Frame::Error(format!("ERR {e}")), + } + } + + #[cfg(feature = "vector")] + Command::VSim { + key, + query, + count, + ef_search, + with_scores, + } => { + let req = ShardRequest::VSim { + key: key.clone(), + query, + count, + ef_search, + }; + match engine.route(&key, req).await { + Ok(ShardResponse::VSimResult(results)) => { + let mut frames = Vec::new(); + for (element, distance) in results { + frames.push(Frame::Bulk(Bytes::from(element))); + if with_scores { + frames.push(Frame::Bulk(Bytes::from(distance.to_string()))); + } + } + Frame::Array(frames) + } + Ok(ShardResponse::WrongType) => wrongtype_error(), + Ok(other) => Frame::Error(format!("ERR unexpected shard response: {other:?}")), + Err(e) => Frame::Error(format!("ERR {e}")), + } + } + + #[cfg(feature = "vector")] + Command::VRem { key, element } => { + let req = ShardRequest::VRem { + key: key.clone(), + element, + }; + match engine.route(&key, req).await { + Ok(ShardResponse::Bool(removed)) => Frame::Integer(if removed { 1 } else { 0 }), + Ok(ShardResponse::WrongType) => wrongtype_error(), + Ok(other) => Frame::Error(format!("ERR unexpected shard response: {other:?}")), + Err(e) => Frame::Error(format!("ERR {e}")), + } + } + + #[cfg(feature = "vector")] + Command::VGet { key, element } => { + let req = ShardRequest::VGet { + key: key.clone(), + element, + }; + match engine.route(&key, req).await { + Ok(ShardResponse::VectorData(Some(vector))) => Frame::Array( + vector + .into_iter() + .map(|v| Frame::Bulk(Bytes::from(v.to_string()))) + .collect(), + ), + Ok(ShardResponse::VectorData(None)) => Frame::Null, + Ok(ShardResponse::WrongType) => wrongtype_error(), + Ok(other) => Frame::Error(format!("ERR unexpected shard response: {other:?}")), + Err(e) => Frame::Error(format!("ERR {e}")), + } + } + + #[cfg(feature = "vector")] + Command::VCard { key } => { + let req = ShardRequest::VCard { key: key.clone() }; + match engine.route(&key, req).await { + Ok(ShardResponse::Integer(count)) => Frame::Integer(count), + Ok(ShardResponse::WrongType) => wrongtype_error(), + Ok(other) => Frame::Error(format!("ERR unexpected shard response: {other:?}")), + Err(e) => Frame::Error(format!("ERR {e}")), + } + } + + #[cfg(feature = "vector")] + Command::VDim { key } => { + let req = ShardRequest::VDim { key: key.clone() }; + match engine.route(&key, req).await { + Ok(ShardResponse::Integer(dim)) => Frame::Integer(dim), + Ok(ShardResponse::WrongType) => wrongtype_error(), + Ok(other) => Frame::Error(format!("ERR unexpected shard response: {other:?}")), + Err(e) => Frame::Error(format!("ERR {e}")), + } + } + + #[cfg(feature = "vector")] + Command::VInfo { key } => { + let req = ShardRequest::VInfo { key: key.clone() }; + match engine.route(&key, req).await { + Ok(ShardResponse::VectorInfo(Some(fields))) => { + let mut frames = Vec::with_capacity(fields.len() * 2); + for (k, v) in fields { + frames.push(Frame::Bulk(Bytes::from(k))); + frames.push(Frame::Bulk(Bytes::from(v))); + } + Frame::Array(frames) + } + Ok(ShardResponse::VectorInfo(None)) => Frame::Null, + Ok(ShardResponse::WrongType) => wrongtype_error(), + Ok(other) => Frame::Error(format!("ERR unexpected shard response: {other:?}")), + Err(e) => Frame::Error(format!("ERR {e}")), + } + } + + #[cfg(not(feature = "vector"))] + Command::VAdd { .. } + | Command::VSim { .. } + | Command::VRem { .. } + | Command::VGet { .. } + | Command::VCard { .. } + | Command::VDim { .. } + | Command::VInfo { .. } => { + Frame::Error("ERR unknown command (vector support not compiled)".into()) + } + #[cfg(feature = "protobuf")] Command::ProtoRegister { name, descriptor } => { let registry = match engine.schema_registry() {