use std::sync::Arc;
use super::{
vector_manager::{
MAX_EXPLORATION_FACTOR, MAX_FILTERING_SCALE_FACTOR, MAX_RETRIEVE_COUNT, MAX_VECTOR_DIMENSIONS,
VectorManager, VectorManagerResult,
},
vector_manager__index::Index,
vector_manager__locking::CreateIndexParams,
vector_types::{VectorDistanceMetricType, VectorQuantType, VectorValueType},
};
use crate::{
objects::types::object_output::ObjectOutput,
resp::{
cmd_strings::GENERIC_ERR_WRONG_NUM_ARGS,
parser::session_parse_state::{strict_f32, strict_i32},
},
storage::session::common::array_key_iteration_functions::cluster_slot,
};
const MIN_M: i32 = 4;
const MAX_M: i32 = 4_096;
const DEFAULT_VSIM_COUNT: i32 = 10;
const DEFAULT_VSIM_EF: i32 = 100;
const DEFAULT_VSIM_FILTER_EF: i32 = 16;
const ERR_VECTOR_SET_DISABLED: &[u8] = b"ERR Vector Set (preview) commands are not enabled";
const ERR_VECTOR_SET_WRONG_TYPE: &[u8] =
b"WRONGTYPE Operation against a key holding the wrong kind of value";
const ERR_INVALID_VECTOR_SPEC: &[u8] = b"ERR invalid vector specification";
const ERR_REDUCE_MUST_BE_POSITIVE: &[u8] = b"REDUCE dimension must be > 0";
const ERR_REDUCE_EXCEEDS_DIMS: &[u8] = b"ERR REDUCE dimension must be <= vector dimensions";
const ERR_INVALID_OPTION_AFTER_ELEMENT: &[u8] = b"ERR invalid option after element";
const ERR_QUANT_SPECIFIED_TWICE: &[u8] = b"Quantization specified multiple times";
const ERR_EF_RANGE: &[u8] = b"ERR EF must be an integer between 1 and 1000000";
const ERR_M_RANGE: &[u8] = b"ERR M must be an integer between 4 and 4096";
const ERR_INVALID_DISTANCE_METRIC: &[u8] = b"ERR invalid XDISTANCE_METRIC";
const ERR_EMPTY_VECTOR_SET_KEY: &[u8] = b"ERR Vector Set key cannot be empty";
const ERR_QUANT_MISMATCH: &[u8] = b"ERR asked quantization mismatch with existing vector set";
const ERR_FP32_MULTIPLE_OF_4: &[u8] = b"FP32 values must be multiple of 4-bytes in size";
const ERR_VALUES_COUNT_MUST_BE_POSITIVE: &[u8] = b"VALUES count must > 0";
const ERR_VALUES_MUST_BE_FLOAT: &[u8] = b"VALUES value must be valid float";
const ERR_VSIM_EXPECTED_KIND: &[u8] = b"VSIM expected ELE, FP32, or VALUES";
const ERR_COUNT_RANGE: &[u8] = b"ERR COUNT must be an integer between 0 and 100000000";
const ERR_EPSILON_MUST_BE_POSITIVE: &[u8] = b"EPSILON must be float > 0";
const ERR_FILTER_EF_RANGE: &[u8] = b"ERR FILTER-EF must be an integer between 4 and 256";
const ERR_UNKNOWN_OPTION: &[u8] = b"Unknown option";
pub(crate) const ERR_ELEMENT_NOT_IN_SET: &[u8] = b"Element not in Vector Set";
const ERR_VEMB_UNEXPECTED_OPTION: &[u8] = b"Unexpected option to VEMB";
const ERR_KEY_NOT_FOUND: &[u8] = b"ERR Key not found";
const ERR_VLINKS_UNEXPECTED_OPTION: &[u8] = b"ERR Unexpected option";
const ERR_EXPECTED_INTEGER_COUNT: &[u8] = b"ERR expected integer count";
const ERR_MAX_ALLOCATIONS_EXCEEDED: &[u8] =
b"ERR Maximum Vector Set allocations exceeded, cannot issue new context";
#[derive(Debug, Clone, PartialEq)]
pub enum VectorReply {
Simple(Vec<u8>),
Error(Vec<u8>),
Integer(i64),
Bulk(Option<Vec<u8>>),
Array(Vec<VectorReply>),
NullArray,
Map(Vec<(VectorReply, VectorReply)>),
Double(f64),
Boolean(bool),
}
impl VectorReply {
pub fn encode_resp2(&self, out: &mut Vec<u8>) {
match self {
VectorReply::Simple(s) => {
out.push(b'+');
out.extend_from_slice(s);
out.extend_from_slice(b"\r\n");
}
VectorReply::Error(e) => {
out.push(b'-');
out.extend_from_slice(e);
out.extend_from_slice(b"\r\n");
}
VectorReply::Integer(i) => {
out.extend_from_slice(format!(":{i}\r\n").as_bytes());
}
VectorReply::Bulk(v) => match v {
Some(v) => {
out.extend_from_slice(format!("${}\r\n", v.len()).as_bytes());
out.extend_from_slice(v);
out.extend_from_slice(b"\r\n");
}
None => out.extend_from_slice(b"$-1\r\n"),
},
VectorReply::NullArray => out.extend_from_slice(b"*-1\r\n"),
VectorReply::Map(pairs) => {
out.extend_from_slice(format!("*{}\r\n", pairs.len() * 2).as_bytes());
for (k, v) in pairs {
k.encode_resp2(out);
v.encode_resp2(out);
}
}
VectorReply::Double(d) => {
let text = ObjectOutput::format_double(*d);
out.extend_from_slice(format!("${}\r\n{}\r\n", text.len(), text).as_bytes());
}
VectorReply::Boolean(b) => {
out.extend_from_slice(if *b { b"$1\r\n1\r\n" } else { b"$0\r\n0\r\n" });
}
VectorReply::Array(items) => {
out.extend_from_slice(format!("*{}\r\n", items.len()).as_bytes());
for item in items {
item.encode_resp2(out);
}
}
}
}
pub fn encode_resp3(&self, out: &mut Vec<u8>) {
match self {
VectorReply::Double(d) => {
out.push(b',');
out.extend_from_slice(ObjectOutput::format_double(*d).as_bytes());
out.extend_from_slice(b"\r\n");
}
VectorReply::Boolean(b) => {
out.extend_from_slice(if *b { b"#t\r\n" } else { b"#f\r\n" });
}
VectorReply::Map(pairs) => {
out.extend_from_slice(format!("%{}\r\n", pairs.len()).as_bytes());
for (k, v) in pairs {
k.encode_resp3(out);
v.encode_resp3(out);
}
}
VectorReply::Array(items) => {
out.extend_from_slice(format!("*{}\r\n", items.len()).as_bytes());
for item in items {
item.encode_resp3(out);
}
}
other => other.encode_resp2(out),
}
}
}
fn eq_ignore_case(a: &[u8], b: &[u8]) -> bool {
a.eq_ignore_ascii_case(b)
}
pub struct RespServerSessionVectors {
pub manager: Arc<VectorManager>,
}
impl RespServerSessionVectors {
pub fn new(manager: Arc<VectorManager>) -> Self {
Self { manager }
}
fn read_index(&self, key: &[u8]) -> Option<Index> {
let stored = self.manager.read_stored_index(key)?;
Index::from_bytes(&stored)
}
pub fn abort_vector_set_wrong_type(&self, key: &[u8]) -> Option<VectorReply> {
if self.manager.read_stored_index(key).is_some() {
Some(VectorReply::Error(ERR_VECTOR_SET_WRONG_TYPE.to_vec()))
} else {
None
}
}
fn abort_disabled(&self) -> VectorReply {
VectorReply::Error(ERR_VECTOR_SET_DISABLED.to_vec())
}
fn abort_wrong_number_of_arguments(cmd: &str) -> VectorReply {
VectorReply::Error(GENERIC_ERR_WRONG_NUM_ARGS.replace("{0}", cmd).into_bytes())
}
pub fn network_vadd(&self, args: &[&[u8]]) -> VectorReply {
self.network_vadd_impl(args, false)
}
pub fn network_vadd_impl(&self, args: &[&[u8]], resp3: bool) -> VectorReply {
if !self.manager.is_enabled {
return self.abort_disabled();
}
if args.len() < 4 {
return Self::abort_wrong_number_of_arguments("VADD");
}
let key = args[0];
let mut cur_ix = 1usize;
let mut reduce_dims = 0u32;
if eq_ignore_case(args[cur_ix], b"REDUCE") {
cur_ix += 1;
let v = args.get(cur_ix).and_then(|a| strict_i32(a));
let Some(v) = v.filter(|v| *v > 0) else {
return VectorReply::Error(ERR_REDUCE_MUST_BE_POSITIVE.to_vec());
};
reduce_dims = v as u32;
cur_ix += 1;
}
let Some(&kind) = args.get(cur_ix) else {
return Self::abort_wrong_number_of_arguments("VADD");
};
let value_type;
let values: Vec<u8>;
let vector_dims: i32;
if eq_ignore_case(kind, b"FP32") {
cur_ix += 1;
let Some(as_bytes) = args.get(cur_ix) else {
return Self::abort_wrong_number_of_arguments("VADD");
};
if as_bytes.len() % 4 != 0 {
return VectorReply::Error(ERR_INVALID_VECTOR_SPEC.to_vec());
}
vector_dims = (as_bytes.len() / 4) as i32;
if vector_dims > MAX_VECTOR_DIMENSIONS as i32 {
return self.abort_too_many_dimensions();
}
value_type = VectorValueType::FP32;
values = as_bytes.to_vec();
cur_ix += 1;
} else if eq_ignore_case(kind, b"VALUES") {
cur_ix += 1;
let Some(count_raw) = args.get(cur_ix) else {
return Self::abort_wrong_number_of_arguments("VADD");
};
let Some(n) = strict_i32(count_raw).filter(|n| *n > 0) else {
return VectorReply::Error(ERR_INVALID_VECTOR_SPEC.to_vec());
};
cur_ix += 1;
if n > MAX_VECTOR_DIMENSIONS as i32 {
return self.abort_too_many_dimensions();
}
if cur_ix + n as usize > args.len() {
return Self::abort_wrong_number_of_arguments("VADD");
}
value_type = VectorValueType::FP32;
let mut floats = Vec::with_capacity(n as usize * 4);
for _ in 0..n {
let Some(f) = args.get(cur_ix).and_then(|a| strict_f32(a, true)) else {
return VectorReply::Error(ERR_INVALID_VECTOR_SPEC.to_vec());
};
floats.extend_from_slice(&f.to_le_bytes());
cur_ix += 1;
}
vector_dims = n;
values = floats;
} else if eq_ignore_case(kind, b"XU8") || eq_ignore_case(kind, b"XB8") {
cur_ix += 1;
let Some(as_bytes) = args.get(cur_ix) else {
return Self::abort_wrong_number_of_arguments("VADD");
};
vector_dims = as_bytes.len() as i32;
if vector_dims > MAX_VECTOR_DIMENSIONS as i32 {
return self.abort_too_many_dimensions();
}
value_type = VectorValueType::XU8;
values = as_bytes.to_vec();
cur_ix += 1;
} else if eq_ignore_case(kind, b"XI8") {
cur_ix += 1;
let Some(as_bytes) = args.get(cur_ix) else {
return Self::abort_wrong_number_of_arguments("VADD");
};
vector_dims = as_bytes.len() as i32;
if vector_dims > MAX_VECTOR_DIMENSIONS as i32 {
return self.abort_too_many_dimensions();
}
value_type = VectorValueType::XI8;
values = as_bytes.to_vec();
cur_ix += 1;
} else {
return VectorReply::Error(ERR_INVALID_VECTOR_SPEC.to_vec());
}
if reduce_dims as i32 > vector_dims {
return VectorReply::Error(ERR_REDUCE_EXCEEDS_DIMS.to_vec());
}
let Some(&element) = args.get(cur_ix) else {
return Self::abort_wrong_number_of_arguments("VADD");
};
cur_ix += 1;
let mut cas_seen = false;
let mut quant: Option<VectorQuantType> = None;
let mut build_ef: Option<i32> = None;
let mut attributes: Option<&[u8]> = None;
let mut num_links: Option<i32> = None;
let mut distance_metric: Option<VectorDistanceMetricType> = None;
while cur_ix < args.len() {
let opt = args[cur_ix];
if eq_ignore_case(opt, b"REDUCE") {
return VectorReply::Error(ERR_INVALID_OPTION_AFTER_ELEMENT.to_vec());
}
if eq_ignore_case(opt, b"CAS") {
if cas_seen {
return VectorReply::Error(b"CAS specified multiple times".to_vec());
}
cas_seen = true;
cur_ix += 1;
} else if eq_ignore_case(opt, b"NOQUANT") {
if quant.is_some() {
return VectorReply::Error(ERR_QUANT_SPECIFIED_TWICE.to_vec());
}
quant = Some(VectorQuantType::NoQuant);
cur_ix += 1;
} else if eq_ignore_case(opt, b"Q8") {
if quant.is_some() {
return VectorReply::Error(ERR_QUANT_SPECIFIED_TWICE.to_vec());
}
quant = Some(VectorQuantType::Q8);
cur_ix += 1;
} else if eq_ignore_case(opt, b"BIN") {
if quant.is_some() {
return VectorReply::Error(ERR_QUANT_SPECIFIED_TWICE.to_vec());
}
quant = Some(VectorQuantType::Bin);
cur_ix += 1;
} else if eq_ignore_case(opt, b"XNOQUANT_U8") || eq_ignore_case(opt, b"XPREQ8") {
if quant.is_some() {
return VectorReply::Error(ERR_QUANT_SPECIFIED_TWICE.to_vec());
}
quant = Some(VectorQuantType::XNoQuant_U8);
cur_ix += 1;
} else if eq_ignore_case(opt, b"XNOQUANT_I8") {
if quant.is_some() {
return VectorReply::Error(ERR_QUANT_SPECIFIED_TWICE.to_vec());
}
quant = Some(VectorQuantType::XNoQuant_I8);
cur_ix += 1;
} else if eq_ignore_case(opt, b"XBIN_I8") {
if quant.is_some() {
return VectorReply::Error(ERR_QUANT_SPECIFIED_TWICE.to_vec());
}
quant = Some(VectorQuantType::XBin_I8);
cur_ix += 1;
} else if eq_ignore_case(opt, b"XBIN_U8") {
if quant.is_some() {
return VectorReply::Error(ERR_QUANT_SPECIFIED_TWICE.to_vec());
}
quant = Some(VectorQuantType::XBin_U8);
cur_ix += 1;
} else if eq_ignore_case(opt, b"EF") {
if build_ef.is_some() {
return VectorReply::Error(b"EF specified multiple times".to_vec());
}
cur_ix += 1;
let Some(v) = args.get(cur_ix) else {
return VectorReply::Error(ERR_INVALID_OPTION_AFTER_ELEMENT.to_vec());
};
let Some(v) = strict_i32(v).filter(|v| *v > 0 && *v <= MAX_EXPLORATION_FACTOR as i32)
else {
return Self::abort_ef_range();
};
build_ef = Some(v);
cur_ix += 1;
} else if eq_ignore_case(opt, b"SETATTR") {
if attributes.is_some() {
return VectorReply::Error(b"SETATTR specified multiple times".to_vec());
}
cur_ix += 1;
let Some(attr) = args.get(cur_ix) else {
return VectorReply::Error(ERR_INVALID_OPTION_AFTER_ELEMENT.to_vec());
};
attributes = Some(attr);
cur_ix += 1;
} else if eq_ignore_case(opt, b"M") {
if num_links.is_some() {
return VectorReply::Error(b"M specified multiple times".to_vec());
}
cur_ix += 1;
let Some(v) = args.get(cur_ix) else {
return VectorReply::Error(ERR_INVALID_OPTION_AFTER_ELEMENT.to_vec());
};
let Some(v) = strict_i32(v).filter(|v| (MIN_M..=MAX_M).contains(v)) else {
return Self::abort_m_range();
};
num_links = Some(v);
cur_ix += 1;
} else if eq_ignore_case(opt, b"XDISTANCE_METRIC") {
if distance_metric.is_some() {
return VectorReply::Error(b"XDISTANCE_METRIC specified multiple times".to_vec());
}
cur_ix += 1;
let Some(metric) = args.get(cur_ix) else {
return VectorReply::Error(ERR_INVALID_OPTION_AFTER_ELEMENT.to_vec());
};
distance_metric = Some(if eq_ignore_case(metric, b"L2") {
VectorDistanceMetricType::L2
} else if eq_ignore_case(metric, b"COSINE") {
VectorDistanceMetricType::Cosine
} else if eq_ignore_case(metric, b"IP") {
VectorDistanceMetricType::InnerProduct
} else if eq_ignore_case(metric, b"XCOSINE_NORMALIZED") {
VectorDistanceMetricType::XCosineNormalized
} else {
return VectorReply::Error(ERR_INVALID_DISTANCE_METRIC.to_vec());
});
cur_ix += 1;
} else {
return VectorReply::Error(ERR_INVALID_OPTION_AFTER_ELEMENT.to_vec());
}
}
if key.is_empty() {
return VectorReply::Error(ERR_EMPTY_VECTOR_SET_KEY.to_vec());
}
let quant = quant.unwrap_or(VectorQuantType::Q8);
let build_ef = build_ef.unwrap_or(200) as u32;
let num_links = num_links.unwrap_or(16) as u32;
let distance_metric = distance_metric.unwrap_or(VectorDistanceMetricType::L2);
if matches!(
quant,
VectorQuantType::XBin_U8
| VectorQuantType::XBin_I8
| VectorQuantType::XNoQuant_U8
| VectorQuantType::XNoQuant_I8
) && reduce_dims != 0
{
return VectorReply::Error(ERR_QUANT_MISMATCH.to_vec());
}
let dims = values.len() as u32
/ match value_type {
VectorValueType::FP32 => 4,
_ => 1,
};
let params = CreateIndexParams {
hash_slot: cluster_slot(key),
dims,
reduce_dims,
quant,
build_exploration_factor: build_ef,
num_links,
distance_metric,
};
let (index, _lock) = match self.manager.read_or_create_vector_index(key, Some(¶ms)) {
Ok(acquired) => acquired,
Err(_) => return VectorReply::Error(ERR_MAX_ALLOCATIONS_EXCEEDED.to_vec()),
};
let stored = index.to_bytes();
match self.manager.try_add(
key,
&stored,
element,
value_type,
&values,
attributes.unwrap_or(b""),
reduce_dims,
quant,
num_links,
distance_metric,
) {
Ok(VectorManagerResult::OK) => {
if resp3 {
VectorReply::Boolean(true)
} else {
VectorReply::Integer(1)
}
}
Ok(VectorManagerResult::Duplicate) => {
if resp3 {
VectorReply::Boolean(false)
} else {
VectorReply::Integer(0)
}
}
Ok(VectorManagerResult::BadParams) => {
VectorReply::Error(VectorManager::error_msg(VectorManagerResult::BadParams).to_vec())
}
Ok(other) => VectorReply::Error(VectorManager::error_msg(other).to_vec()),
Err(e) => VectorReply::Error(e.message),
}
}
fn abort_too_many_dimensions(&self) -> VectorReply {
VectorReply::Error(
format!("ERR vector exceeds maximum of {MAX_VECTOR_DIMENSIONS} dimensions").into_bytes(),
)
}
fn abort_ef_range() -> VectorReply {
VectorReply::Error(ERR_EF_RANGE.to_vec())
}
fn abort_m_range() -> VectorReply {
VectorReply::Error(ERR_M_RANGE.to_vec())
}
pub fn network_vsim(&self, args: &[&[u8]]) -> VectorReply {
self.network_vsim_impl(args, false)
}
pub fn network_vsim_impl(&self, args: &[&[u8]], resp3: bool) -> VectorReply {
if !self.manager.is_enabled {
return self.abort_disabled();
}
if args.len() < 3 {
return Self::abort_wrong_number_of_arguments("VSIM");
}
let key = args[0];
let kind = args[1];
let mut cur_ix = 2usize;
let mut element: Option<&[u8]> = None;
let mut value_type = VectorValueType::Invalid;
let mut values: Vec<u8> = Vec::new();
if eq_ignore_case(kind, b"ELE") {
element = Some(args.get(cur_ix).copied().unwrap_or(b""));
cur_ix += 1;
} else if eq_ignore_case(kind, b"FP32") {
let Some(as_bytes) = args.get(cur_ix) else {
return Self::abort_wrong_number_of_arguments("VSIM");
};
if as_bytes.len() % 4 != 0 {
return VectorReply::Error(ERR_FP32_MULTIPLE_OF_4.to_vec());
}
if as_bytes.len() / 4 > MAX_VECTOR_DIMENSIONS as usize {
return self.abort_too_many_dimensions();
}
value_type = VectorValueType::FP32;
values = as_bytes.to_vec();
cur_ix += 1;
} else if eq_ignore_case(kind, b"XU8") || eq_ignore_case(kind, b"XB8") {
let Some(as_bytes) = args.get(cur_ix) else {
return Self::abort_wrong_number_of_arguments("VSIM");
};
if as_bytes.len() > MAX_VECTOR_DIMENSIONS as usize {
return self.abort_too_many_dimensions();
}
value_type = VectorValueType::XU8;
values = as_bytes.to_vec();
cur_ix += 1;
} else if eq_ignore_case(kind, b"XI8") {
let Some(as_bytes) = args.get(cur_ix) else {
return Self::abort_wrong_number_of_arguments("VSIM");
};
if as_bytes.len() > MAX_VECTOR_DIMENSIONS as usize {
return self.abort_too_many_dimensions();
}
value_type = VectorValueType::XI8;
values = as_bytes.to_vec();
cur_ix += 1;
} else if eq_ignore_case(kind, b"VALUES") {
let Some(count_raw) = args.get(cur_ix) else {
return Self::abort_wrong_number_of_arguments("VSIM");
};
let Some(n) = strict_i32(count_raw).filter(|n| *n > 0) else {
return VectorReply::Error(ERR_VALUES_COUNT_MUST_BE_POSITIVE.to_vec());
};
if n > MAX_VECTOR_DIMENSIONS as i32 {
return self.abort_too_many_dimensions();
}
cur_ix += 1;
if cur_ix + n as usize > args.len() {
return Self::abort_wrong_number_of_arguments("VSIM");
}
value_type = VectorValueType::FP32;
for _ in 0..n {
let Some(f) = args.get(cur_ix).and_then(|a| strict_f32(a, true)) else {
return VectorReply::Error(ERR_VALUES_MUST_BE_FLOAT.to_vec());
};
values.extend_from_slice(&f.to_le_bytes());
cur_ix += 1;
}
} else {
return VectorReply::Error(ERR_VSIM_EXPECTED_KIND.to_vec());
}
let mut with_scores = false;
let mut with_attribs = false;
let mut count: Option<i32> = None;
let mut epsilon: Option<f32> = None;
let mut ef: Option<i32> = None;
let mut filter: Option<&[u8]> = None;
let mut filter_ef: Option<i32> = None;
let mut truth = false;
let mut no_thread = false;
while cur_ix < args.len() {
let opt = args[cur_ix];
if eq_ignore_case(opt, b"WITHSCORES") {
if with_scores {
return VectorReply::Error(b"WITHSCORES specified multiple times".to_vec());
}
with_scores = true;
cur_ix += 1;
} else if eq_ignore_case(opt, b"WITHATTRIBS") {
if with_attribs {
return VectorReply::Error(b"WITHATTRIBS specified multiple times".to_vec());
}
with_attribs = true;
cur_ix += 1;
} else if eq_ignore_case(opt, b"COUNT") {
if count.is_some() {
return VectorReply::Error(b"COUNT specified multiple times".to_vec());
}
cur_ix += 1;
let Some(v) = args.get(cur_ix) else {
return Self::abort_wrong_number_of_arguments("VSIM");
};
let Some(v) = strict_i32(v).filter(|v| *v >= 0 && *v <= MAX_RETRIEVE_COUNT as i32) else {
return Self::abort_count_range();
};
count = Some(v);
cur_ix += 1;
} else if eq_ignore_case(opt, b"EPSILON") {
if epsilon.is_some() {
return VectorReply::Error(b"EPSILON specified multiple times".to_vec());
}
cur_ix += 1;
let Some(v) = args.get(cur_ix) else {
return Self::abort_wrong_number_of_arguments("VSIM");
};
let Some(v) = strict_f32(v, true).filter(|v| *v > 0.0) else {
return VectorReply::Error(ERR_EPSILON_MUST_BE_POSITIVE.to_vec());
};
epsilon = Some(v);
cur_ix += 1;
} else if eq_ignore_case(opt, b"EF") {
if ef.is_some() {
return VectorReply::Error(b"EF specified multiple times".to_vec());
}
cur_ix += 1;
let Some(v) = args.get(cur_ix) else {
return Self::abort_wrong_number_of_arguments("VSIM");
};
let Some(v) = strict_i32(v).filter(|v| *v > 0 && *v <= MAX_EXPLORATION_FACTOR as i32)
else {
return Self::abort_ef_range();
};
ef = Some(v);
cur_ix += 1;
} else if eq_ignore_case(opt, b"FILTER") {
if filter.is_some() {
return VectorReply::Error(b"FILTER specified multiple times".to_vec());
}
cur_ix += 1;
let Some(f) = args.get(cur_ix) else {
return Self::abort_wrong_number_of_arguments("VSIM");
};
filter = Some(f);
cur_ix += 1;
} else if eq_ignore_case(opt, b"FILTER-EF") {
if filter_ef.is_some() {
return VectorReply::Error(b"FILTER-EF specified multiple times".to_vec());
}
cur_ix += 1;
let Some(v) = args.get(cur_ix) else {
return Self::abort_wrong_number_of_arguments("VSIM");
};
let Some(v) = strict_i32(v).filter(|v| *v >= 4 && *v <= MAX_FILTERING_SCALE_FACTOR as i32)
else {
return Self::abort_filter_ef_range();
};
filter_ef = Some(v);
cur_ix += 1;
} else if eq_ignore_case(opt, b"TRUTH") {
if truth {
return VectorReply::Error(b"TRUTH specified multiple times".to_vec());
}
truth = true;
cur_ix += 1;
} else if eq_ignore_case(opt, b"NOTHREAD") {
if no_thread {
return VectorReply::Error(b"NOTHREAD specified multiple times".to_vec());
}
no_thread = true;
cur_ix += 1;
} else {
return VectorReply::Error(ERR_UNKNOWN_OPTION.to_vec());
}
}
let _ = (truth, no_thread);
let delta = epsilon.unwrap_or(f32::INFINITY);
let filter_effort = filter_ef.unwrap_or(DEFAULT_VSIM_FILTER_EF).max(0) as usize;
let count = count.unwrap_or(DEFAULT_VSIM_COUNT);
let Some(stored) = self.manager.read_stored_index(key) else {
return VectorReply::Array(Vec::new());
};
let result = match element {
Some(elem) => self.manager.element_similarity(
&stored,
elem,
count.max(0) as usize,
ef.unwrap_or(DEFAULT_VSIM_EF).max(0) as usize,
filter.unwrap_or(b""),
filter_effort,
delta,
with_attribs,
),
None => self.manager.value_similarity(
&stored,
value_type,
&values,
count.max(0) as usize,
ef.unwrap_or(DEFAULT_VSIM_EF).max(0) as usize,
filter.unwrap_or(b""),
filter_effort,
delta,
with_attribs,
),
};
let output = match result {
Ok(out) => out,
Err(e) => return VectorReply::Error(e.message),
};
let ids: Vec<&[u8]> = super::vector_manager::unpack_length_prefixed(&output.output_ids);
let attrs: Vec<&[u8]> =
super::vector_manager::unpack_length_prefixed(&output.output_attributes);
if resp3 {
Self::write_resp3_result(
count as usize,
&ids,
&output.output_distances,
&output.filter_bitmap,
with_attribs.then_some(&attrs),
with_scores,
)
} else {
Self::write_resp2_result(
count as usize,
&ids,
&output.output_distances,
&output.filter_bitmap,
with_attribs.then_some(&attrs),
with_scores,
)
}
}
pub fn write_resp3_result(
count: usize,
ids: &[&[u8]],
distances: &[f32],
filter_bitmap: &[u8],
attributes: Option<&Vec<&[u8]>>,
with_scores: bool,
) -> VectorReply {
let has_filter = !filter_bitmap.is_empty();
let with_attribs = attributes.is_some();
let total_found = ids.len();
let output_count = if has_filter {
filter_bitmap
.iter()
.map(|b| b.count_ones() as usize)
.sum::<usize>()
.min(count)
} else {
total_found.min(count)
};
let mut plain = Vec::new();
let mut map = Vec::new();
let mut written = 0usize;
for (result_index, id) in ids.iter().enumerate().take(total_found) {
if written >= output_count {
break;
}
if has_filter && (filter_bitmap[result_index >> 3] >> (result_index & 7)) & 1 == 0 {
continue;
}
let score = VectorReply::Double(f64::from(
distances.get(result_index).copied().unwrap_or(0.0),
));
let attr_reply = match attributes
.and_then(|attrs| attrs.get(result_index))
.copied()
{
Some(a) if !a.is_empty() => VectorReply::Bulk(Some(a.to_vec())),
_ => VectorReply::Bulk(None),
};
if !with_scores && !with_attribs {
plain.push(VectorReply::Bulk(Some(id.to_vec())));
} else {
let value = if with_scores && with_attribs {
VectorReply::Array(vec![score, attr_reply])
} else if with_scores {
score
} else {
attr_reply
};
map.push((VectorReply::Bulk(Some(id.to_vec())), value));
}
written += 1;
}
if !with_scores && !with_attribs {
VectorReply::Array(plain)
} else {
VectorReply::Map(map)
}
}
pub fn write_resp2_result(
count: usize,
ids: &[&[u8]],
distances: &[f32],
filter_bitmap: &[u8],
attributes: Option<&Vec<&[u8]>>,
with_scores: bool,
) -> VectorReply {
let has_filter = !filter_bitmap.is_empty();
let with_attribs = attributes.is_some();
let total_found = ids.len();
let output_count = if has_filter {
filter_bitmap
.iter()
.map(|b| b.count_ones() as usize)
.sum::<usize>()
.min(count)
} else {
total_found.min(count)
};
let mut items = Vec::new();
let mut written = 0usize;
for (result_index, id) in ids.iter().enumerate().take(total_found) {
if written >= output_count {
break;
}
if has_filter && (filter_bitmap[result_index >> 3] >> (result_index & 7)) & 1 == 0 {
continue;
}
items.push(VectorReply::Bulk(Some(id.to_vec())));
if with_scores {
items.push(VectorReply::Double(f64::from(
distances.get(result_index).copied().unwrap_or(0.0),
)));
}
if with_attribs && let Some(attr) = attributes.and_then(|attrs| attrs.get(result_index)) {
items.push(VectorReply::Bulk(Some(attr.to_vec())));
}
written += 1;
}
VectorReply::Array(items)
}
pub fn network_vemb(&self, args: &[&[u8]]) -> VectorReply {
if !self.manager.is_enabled {
return self.abort_disabled();
}
if args.len() < 2 || args.len() > 3 {
return Self::abort_wrong_number_of_arguments("VEMB");
}
let raw = if args.len() == 3 {
if !eq_ignore_case(args[2], b"RAW") {
return VectorReply::Error(ERR_VEMB_UNEXPECTED_OPTION.to_vec());
}
true
} else {
false
};
let Some(stored) = self.manager.read_stored_index(args[0]) else {
return VectorReply::Array(Vec::new());
};
if raw {
return match self.manager.try_get_raw_embedding(&stored, args[1]) {
Some((bytes, quant, norm, range)) => {
let quant_name: &[u8] = match quant {
VectorQuantType::Bin | VectorQuantType::XBin_I8 | VectorQuantType::XBin_U8 => b"bin",
VectorQuantType::Q8 | VectorQuantType::XNoQuant_U8 | VectorQuantType::XNoQuant_I8 => {
b"q8"
}
VectorQuantType::NoQuant => b"fp32",
VectorQuantType::Invalid => b"fp32",
};
let mut items = vec![
VectorReply::Simple(quant_name.to_vec()),
VectorReply::Bulk(Some(bytes)),
VectorReply::Double(norm),
];
if quant == VectorQuantType::Q8 {
items.push(VectorReply::Double(range.unwrap_or(0.0)));
}
VectorReply::Array(items)
}
None => VectorReply::Array(Vec::new()),
};
}
match self.manager.try_get_embedding(&stored, args[1]) {
Some(embedding) => VectorReply::Array(
embedding
.into_iter()
.map(|v| VectorReply::Double(f64::from(v)))
.collect(),
),
None => VectorReply::Array(Vec::new()),
}
}
pub fn network_vcard(&self, args: &[&[u8]]) -> VectorReply {
if !self.manager.is_enabled {
return self.abort_disabled();
}
if args.len() != 1 {
return Self::abort_wrong_number_of_arguments("VCARD");
}
let Some(index) = self.read_index(args[0]) else {
return VectorReply::Integer(0);
};
VectorReply::Integer(self.manager.service.card(index.context) as i64)
}
pub fn network_vdim(&self, args: &[&[u8]]) -> VectorReply {
if !self.manager.is_enabled {
return self.abort_disabled();
}
if args.len() != 1 {
return Self::abort_wrong_number_of_arguments("VDIM");
}
let Some(index) = self.read_index(args[0]) else {
return VectorReply::Error(ERR_KEY_NOT_FOUND.to_vec());
};
VectorReply::Integer(i64::from(index.dimensions))
}
pub fn network_vgetattr(&self, args: &[&[u8]]) -> VectorReply {
if !self.manager.is_enabled {
return self.abort_disabled();
}
if args.len() != 2 {
return Self::abort_wrong_number_of_arguments("VGETATTR");
}
let Some(index) = self.read_index(args[0]) else {
return VectorReply::Bulk(None);
};
match self.manager.service.get_attribute(index.context, args[1]) {
Some(attr) => VectorReply::Bulk(Some(attr)),
None => VectorReply::Bulk(None),
}
}
pub fn network_vinfo(&self, args: &[&[u8]]) -> VectorReply {
if !self.manager.is_enabled {
return self.abort_disabled();
}
if args.len() != 1 {
return Self::abort_wrong_number_of_arguments("VINFO");
}
let Some(index) = self.read_index(args[0]) else {
return VectorReply::NullArray;
};
let quant: &[u8] = match index.quant_type {
VectorQuantType::NoQuant => b"f32",
VectorQuantType::Bin => b"bin",
VectorQuantType::Q8 => b"q8",
VectorQuantType::XNoQuant_U8 => b"xnoquant_u8",
VectorQuantType::XNoQuant_I8 => b"xnoquant_i8",
VectorQuantType::XBin_I8 => b"xbin_i8",
VectorQuantType::XBin_U8 => b"xbin_u8",
VectorQuantType::Invalid => {
return VectorReply::Error(b"ERR Invalid VectorQuantType".to_vec());
}
};
let metric: &[u8] = match index.distance_metric {
VectorDistanceMetricType::Cosine => b"cosine",
VectorDistanceMetricType::InnerProduct => b"inner-product",
VectorDistanceMetricType::L2 => b"l2",
VectorDistanceMetricType::XCosineNormalized => b"cosine-normalized",
};
let bulk_u32 = |v: u32| VectorReply::Bulk(Some(v.to_string().into_bytes()));
VectorReply::Array(vec![
VectorReply::Simple(b"quant-type".to_vec()),
VectorReply::Simple(quant.to_vec()),
VectorReply::Simple(b"distance-metric".to_vec()),
VectorReply::Simple(metric.to_vec()),
VectorReply::Simple(b"input-vector-dimensions".to_vec()),
bulk_u32(index.dimensions),
VectorReply::Simple(b"reduced-dimensions".to_vec()),
bulk_u32(index.reduce_dims),
VectorReply::Simple(b"build-exploration-factor".to_vec()),
bulk_u32(index.build_exploration_factor),
VectorReply::Simple(b"num-links".to_vec()),
bulk_u32(index.num_links),
VectorReply::Simple(b"size".to_vec()),
VectorReply::Bulk(Some(
self
.manager
.service
.card(index.context)
.to_string()
.into_bytes(),
)),
])
}
pub fn network_vismember(&self, args: &[&[u8]]) -> VectorReply {
self.network_vismember_impl(args, false)
}
pub fn network_vismember_impl(&self, args: &[&[u8]], resp3: bool) -> VectorReply {
if !self.manager.is_enabled {
return self.abort_disabled();
}
if args.len() != 2 {
return Self::abort_wrong_number_of_arguments("VISMEMBER");
}
let member = self
.read_index(args[0])
.is_some_and(|index| self.manager.is_member(&index.to_bytes(), args[1]));
match (member, resp3) {
(true, true) => VectorReply::Boolean(true),
(false, true) => VectorReply::Boolean(false),
(true, false) => VectorReply::Integer(1),
(false, false) => VectorReply::Integer(0),
}
}
pub fn network_vlinks(&self, args: &[&[u8]]) -> VectorReply {
if !self.manager.is_enabled {
return self.abort_disabled();
}
if args.len() != 2 && args.len() != 3 {
return Self::abort_wrong_number_of_arguments("VLINKS");
}
if args.len() == 3 && !eq_ignore_case(args[2], b"WITHSCORES") {
return VectorReply::Error(ERR_VLINKS_UNEXPECTED_OPTION.to_vec());
}
let Some(index) = self.read_index(args[0]) else {
return VectorReply::Bulk(None);
};
match self.manager.service.links_of(index.context, args[1]) {
Some(links) => VectorReply::Array(
links
.into_iter()
.map(|l| VectorReply::Bulk(Some(l)))
.collect(),
),
None => VectorReply::Bulk(None),
}
}
pub fn network_vrandmember(&self, args: &[&[u8]]) -> VectorReply {
if !self.manager.is_enabled {
return self.abort_disabled();
}
if args.is_empty() || args.len() > 2 {
return Self::abort_wrong_number_of_arguments("VRANDMEMBER");
}
let count = match args.get(1) {
Some(raw) => {
let Some(v) = strict_i32(raw) else {
return VectorReply::Error(ERR_EXPECTED_INTEGER_COUNT.to_vec());
};
v
}
None => 1,
};
let Some(index) = self.read_index(args[0]) else {
return if args.len() == 2 {
VectorReply::Array(Vec::new())
} else {
VectorReply::Bulk(None)
};
};
let samples = self
.manager
.service
.sample(index.context, count.max(0) as usize);
if count == 1 && args.len() == 1 {
return match samples.into_iter().next() {
Some(s) => VectorReply::Bulk(Some(s)),
None => VectorReply::Bulk(None),
};
}
VectorReply::Array(
samples
.into_iter()
.map(|v| VectorReply::Bulk(Some(v)))
.collect(),
)
}
pub fn network_vrem(&self, args: &[&[u8]]) -> VectorReply {
if !self.manager.is_enabled {
return self.abort_disabled();
}
if args.len() != 2 {
return Self::abort_wrong_number_of_arguments("VREM");
}
let Some(index) = self.read_index(args[0]) else {
return VectorReply::Integer(0);
};
let removed = self.manager.try_remove(&index.to_bytes(), args[1]);
VectorReply::Integer(match removed {
VectorManagerResult::OK => 1,
_ => 0,
})
}
pub fn network_vsetattr(&self, args: &[&[u8]]) -> VectorReply {
self.network_vsetattr_impl(args, false)
}
pub fn network_vsetattr_impl(&self, args: &[&[u8]], resp3: bool) -> VectorReply {
if !self.manager.is_enabled {
return self.abort_disabled();
}
if args.len() != 3 {
return Self::abort_wrong_number_of_arguments("VSETATTR");
}
let found = self.read_index(args[0]).is_some_and(|index| {
self
.manager
.try_set_attribute(&index.to_bytes(), args[1], args[2])
});
match (found, resp3) {
(true, true) => VectorReply::Boolean(true),
(false, true) => VectorReply::Boolean(false),
(true, false) => VectorReply::Integer(1),
(false, false) => VectorReply::Integer(0),
}
}
}
impl RespServerSessionVectors {
fn abort_count_range() -> VectorReply {
VectorReply::Error(ERR_COUNT_RANGE.to_vec())
}
fn abort_filter_ef_range() -> VectorReply {
VectorReply::Error(ERR_FILTER_EF_RANGE.to_vec())
}
}
#[cfg(test)]
mod tests {
use std::str::from_utf8;
use super::{
super::{
vector_manager::{VectorManager, VectorManagerOptions},
vector_types::{VectorDistanceMetricType, VectorIdFormat, VectorQuantType},
},
*,
};
fn session() -> RespServerSessionVectors {
let manager = Arc::new(VectorManager::new(VectorManagerOptions {
is_enabled: true,
..Default::default()
}));
RespServerSessionVectors::new(manager)
}
fn s(bytes: &[u8]) -> &str {
from_utf8(bytes).unwrap()
}
fn err_text(r: VectorReply) -> String {
match r {
VectorReply::Error(e) => s(&e).to_owned(),
other => panic!("期望错误应答,实际 {other:?}"),
}
}
#[test]
fn vadd_parse_and_defaults() {
let sess = session();
let r = sess.network_vadd(&[b"k", b"FP32"]);
assert!(err_text(r).contains("wrong number of arguments"));
let r = sess.network_vadd(&[b"k", b"VALUES", b"2", b"1.5", b"-2.5", b"elem1", b"CAS"]);
assert_eq!(r, VectorReply::Integer(1));
let index = sess.read_index(b"k").unwrap();
assert_eq!(index.dimensions, 2);
assert_eq!(index.quant_type, VectorQuantType::Q8);
assert_eq!(index.num_links, 16);
assert_eq!(index.build_exploration_factor, 200);
assert_eq!(index.distance_metric, VectorDistanceMetricType::L2);
assert!(sess.manager.is_member(&index.to_bytes(), b"elem1"));
let r = sess.network_vadd(&[b"k", b"VALUES", b"2", b"1.5", b"-2.5", b"elem1"]);
assert_eq!(r, VectorReply::Integer(0));
let r = sess.network_vadd(&[
b"k2",
b"FP32",
&f32_bytes(&[3.0, 4.0]),
b"elem2",
b"NOQUANT",
b"EF",
b"64",
b"SETATTR",
b"{\"a\":1}",
b"M",
b"8",
b"XDISTANCE_METRIC",
b"COSINE",
]);
assert_eq!(r, VectorReply::Integer(1));
let index = sess.read_index(b"k2").unwrap();
assert_eq!(index.quant_type, VectorQuantType::NoQuant);
assert_eq!(index.num_links, 8);
assert_eq!(index.build_exploration_factor, 64);
assert_eq!(index.distance_metric, VectorDistanceMetricType::Cosine);
assert_eq!(
sess.manager.service.get_attribute(index.context, b"elem2"),
Some(b"{\"a\":1}".to_vec())
);
let r = sess.network_vadd(&[
b"kr",
b"REDUCE",
b"1",
b"FP32",
&f32_bytes(&[3.0, 4.0]),
b"e0",
]);
assert_eq!(r, VectorReply::Integer(1));
let index = sess.read_index(b"kr").unwrap();
assert_eq!(index.reduce_dims, 1);
let r = sess.network_vadd(&[
b"kr2",
b"REDUCE",
b"5",
b"FP32",
&f32_bytes(&[1.0, 2.0]),
b"e",
]);
assert_eq!(
err_text(r),
"ERR REDUCE dimension must be <= vector dimensions"
);
let r = sess.network_vadd(&[
b"kq",
b"REDUCE",
b"1",
b"FP32",
&f32_bytes(&[1.0, 2.0]),
b"e",
b"XNOQUANT_U8",
]);
assert_eq!(
err_text(r),
"ERR asked quantization mismatch with existing vector set"
);
let v1 = f32_bytes(&[1.0]);
let dup_sets: Vec<Vec<&[u8]>> = vec![
vec![b"k", b"FP32", &v1, b"x", b"NOQUANT", b"Q8"],
vec![b"k", b"FP32", &v1, b"x", b"EF", b"8", b"EF", b"9"],
vec![b"k", b"FP32", &v1, b"x", b"M", b"8", b"M", b"9"],
vec![b"k", b"FP32", &v1, b"x", b"SETATTR", b"a", b"SETATTR", b"b"],
vec![b"k", b"FP32", &v1, b"x", b"CAS", b"CAS"],
vec![
b"k",
b"FP32",
&v1,
b"x",
b"XDISTANCE_METRIC",
b"L2",
b"XDISTANCE_METRIC",
b"COSINE",
],
];
for dup in dup_sets {
let r = sess.network_vadd(&dup);
assert!(matches!(r, VectorReply::Error(_)), "重复选项应报错: {r:?}");
}
assert_eq!(
err_text(sess.network_vadd(&[b"k", b"FP32", &v1, b"x", b"NOQUANT", b"Q8"])),
"Quantization specified multiple times"
);
let r = sess.network_vadd(&[b"k", b"FP32", &f32_bytes(&[1.0]), b"x", b"M", b"2"]);
assert_eq!(err_text(r), "ERR M must be an integer between 4 and 4096");
assert_eq!(
err_text(sess.network_vadd(&[b"k", b"FP32", &v1, b"x", b"EF", b"0"])),
"ERR EF must be an integer between 1 and 1000000"
);
assert_eq!(
err_text(sess.network_vadd(&[b"k", b"FP32", &v1, b"x", b"XDISTANCE_METRIC", b"DOT"])),
"ERR invalid XDISTANCE_METRIC"
);
assert_eq!(
err_text(sess.network_vadd(&[b"k", b"FP32", &v1, b"x", b"WHAT"])),
"ERR invalid option after element"
);
assert_eq!(
err_text(sess.network_vadd(&[b"k", b"FP32", &v1, b"x", b"REDUCE", b"2"])),
"ERR invalid option after element"
);
assert_eq!(
err_text(sess.network_vadd(&[b"k", b"FP32", b"123", b"x"])),
"ERR invalid vector specification"
);
assert_eq!(
err_text(sess.network_vadd(&[b"", b"FP32", &v1, b"x"])),
"ERR Vector Set key cannot be empty"
);
let r = sess.network_vadd(&[b"k", b"VALUES", b"02", b"1.0", b"2.0", b"e"]);
assert_eq!(err_text(r), "ERR invalid vector specification");
}
#[test]
fn vsim_options_and_output() {
let sess = session();
sess
.manager
.try_add(
b"vs",
&seed_index(&sess, b"vs", 2),
b"near",
VectorValueType::FP32,
&f32_bytes(&[1.0, 0.0]),
b"{\"n\":1}",
0,
VectorQuantType::NoQuant,
8,
VectorDistanceMetricType::L2,
)
.unwrap();
sess
.manager
.try_add(
b"vs",
&sess.manager.read_stored_index(b"vs").unwrap(),
b"far",
VectorValueType::FP32,
&f32_bytes(&[9.0, 9.0]),
b"{\"n\":2}",
0,
VectorQuantType::NoQuant,
8,
VectorDistanceMetricType::L2,
)
.unwrap();
let r = sess.network_vsim(&[
b"vs",
b"FP32",
&f32_bytes(&[1.0, 0.0]),
b"WITHSCORES",
b"COUNT",
b"2",
]);
let mut encoded = Vec::new();
r.encode_resp2(&mut encoded);
let text = s(&encoded);
assert!(text.starts_with("*4\r\n"), "扁平数组 id/score 成对: {text}");
let r3 = sess.network_vsim_impl(
&[
b"vs",
b"FP32",
&f32_bytes(&[1.0, 0.0]),
b"WITHSCORES",
b"COUNT",
b"2",
],
true,
);
let mut encoded3 = Vec::new();
r3.encode_resp3(&mut encoded3);
let text3 = s(&encoded3);
assert!(text3.starts_with("%2\r\n"), "RESP3 map 头: {text3}");
let r3 = sess.network_vsim_impl(&[b"vs", b"ELE", b"near", b"COUNT", b"1"], true);
let mut encoded3 = Vec::new();
r3.encode_resp3(&mut encoded3);
assert!(
s(&encoded3).starts_with("*1\r\n"),
"RESP3 普通数组: {}",
s(&encoded3)
);
let r = sess.network_vsim(&[
b"vs",
b"FP32",
&f32_bytes(&[1.0, 0.0]),
b"FILTER",
b".n > 1",
b"COUNT",
b"2",
]);
let mut encoded = Vec::new();
r.encode_resp2(&mut encoded);
let text = s(&encoded);
assert!(text.contains("far") && !text.contains("near"));
let r = sess.network_vsim(&[b"vs", b"ELE", b"near", b"COUNT", b"1"]);
let mut encoded = Vec::new();
r.encode_resp2(&mut encoded);
assert!(s(&encoded).contains("near"));
let r = sess.network_vsim(&[
b"vs",
b"FP32",
&f32_bytes(&[1.0, 0.0]),
b"COUNT",
b"1",
b"COUNT",
b"2",
]);
assert_eq!(err_text(r), "COUNT specified multiple times");
assert_eq!(
err_text(sess.network_vsim(&[b"vs", b"ELE", b"near", b"WHAT"])),
"Unknown option"
);
assert_eq!(
err_text(sess.network_vsim(&[b"vs", b"WHAT", b"x"])),
"VSIM expected ELE, FP32, or VALUES"
);
assert_eq!(
err_text(sess.network_vsim(&[b"vs", b"VALUES", b"2", b"1.0", b"abc"])),
"VALUES value must be valid float"
);
assert!(
err_text(sess.network_vsim(&[b"vs", b"ELE", b"near", b"COUNT"]))
.contains("wrong number of arguments")
);
assert_eq!(
err_text(sess.network_vsim(&[b"vs", b"ELE", b"ghost"])),
"Element not in Vector Set"
);
assert_eq!(
sess.network_vsim(&[b"nope", b"ELE", b"x"]),
VectorReply::Array(Vec::new())
);
assert_eq!(
err_text(sess.network_vsim(&[b"vs", b"ELE", b"near", b"EPSILON", b"-1"])),
"EPSILON must be float > 0"
);
assert_eq!(
err_text(sess.network_vsim(&[b"vs", b"ELE", b"near", b"FILTER-EF", b"3"])),
"ERR FILTER-EF must be an integer between 4 and 256"
);
}
#[test]
fn auxiliary_commands() {
let sess = session();
sess
.manager
.try_add(
b"aux",
&seed_index(&sess, b"aux", 2),
b"e1",
VectorValueType::FP32,
&f32_bytes(&[5.0, 6.0]),
b"{\"tag\":\"x\"}",
0,
VectorQuantType::NoQuant,
8,
VectorDistanceMetricType::L2,
)
.unwrap();
sess
.manager
.try_add(
b"aux",
&sess.manager.read_stored_index(b"aux").unwrap(),
b"e2",
VectorValueType::FP32,
&f32_bytes(&[1.0, 1.0]),
b"",
0,
VectorQuantType::NoQuant,
8,
VectorDistanceMetricType::L2,
)
.unwrap();
assert_eq!(sess.network_vcard(&[b"aux"]), VectorReply::Integer(2));
assert_eq!(sess.network_vdim(&[b"aux"]), VectorReply::Integer(2));
assert_eq!(
sess.network_vismember(&[b"aux", b"e1"]),
VectorReply::Integer(1)
);
assert_eq!(
sess.network_vismember(&[b"aux", b"zz"]),
VectorReply::Integer(0)
);
assert_eq!(
sess.network_vismember_impl(&[b"aux", b"e1"], true),
VectorReply::Boolean(true)
);
let emb = sess.network_vemb(&[b"aux", b"e1"]);
let mut encoded = Vec::new();
emb.encode_resp2(&mut encoded);
let text = s(&encoded);
assert!(text.contains('5'), "嵌入值包含 5: {text}");
let raw = sess.network_vemb(&[b"aux", b"e1", b"RAW"]);
let VectorReply::Array(items) = raw else {
panic!("RAW 应为数组: {raw:?}");
};
assert_eq!(items.len(), 3);
assert_eq!(items[0], VectorReply::Simple(b"fp32".to_vec()));
assert_eq!(items[1], VectorReply::Bulk(Some(f32_bytes(&[5.0, 6.0]))));
assert_eq!(
err_text(sess.network_vemb(&[b"aux", b"e1", b"BAD"])),
"Unexpected option to VEMB"
);
assert_eq!(
sess.network_vemb(&[b"aux", b"ghost"]),
VectorReply::Array(Vec::new())
);
assert_eq!(
sess.network_vgetattr(&[b"aux", b"e1"]),
VectorReply::Bulk(Some(b"{\"tag\":\"x\"}".to_vec()))
);
assert_eq!(
sess.network_vsetattr(&[b"aux", b"e1", b"{}"]),
VectorReply::Integer(1)
);
assert_eq!(
sess.network_vsetattr(&[b"aux", b"ghost", b"{}"]),
VectorReply::Integer(0)
);
assert_eq!(
sess.network_vgetattr(&[b"aux", b"e1"]),
VectorReply::Bulk(Some(b"{}".to_vec()))
);
assert!(matches!(
sess.network_vlinks(&[b"aux", b"e1"]),
VectorReply::Array(_)
));
assert_eq!(
sess.network_vlinks(&[b"aux", b"ghost"]),
VectorReply::Bulk(None)
);
assert_eq!(
err_text(sess.network_vlinks(&[b"aux", b"e1", b"BAD"])),
"ERR Unexpected option"
);
assert!(matches!(
sess.network_vrandmember(&[b"aux", b"2"]),
VectorReply::Array(_)
));
assert_eq!(
err_text(sess.network_vrandmember(&[b"aux", b"x"])),
"ERR expected integer count"
);
assert_eq!(
sess.network_vrandmember(&[b"ghost", b"2"]),
VectorReply::Array(Vec::new())
);
assert_eq!(
sess.network_vrandmember(&[b"ghost"]),
VectorReply::Bulk(None)
);
let info = sess.network_vinfo(&[b"aux"]);
let mut encoded = Vec::new();
info.encode_resp2(&mut encoded);
let text = s(&encoded);
assert!(text.starts_with("*14\r\n"), "VINFO 14 项: {text}");
assert!(text.contains("input-vector-dimensions"));
assert!(text.contains("reduced-dimensions"));
assert!(text.contains("f32"));
assert!(text.contains("l2"));
assert_eq!(sess.network_vinfo(&[b"ghost"]), VectorReply::NullArray);
assert_eq!(sess.network_vrem(&[b"aux", b"e2"]), VectorReply::Integer(1));
assert_eq!(sess.network_vrem(&[b"aux", b"e2"]), VectorReply::Integer(0));
assert_eq!(sess.network_vcard(&[b"aux"]), VectorReply::Integer(1));
assert_eq!(
err_text(sess.network_vdim(&[b"ghost"])),
"ERR Key not found"
);
}
#[test]
fn disabled_and_reply_encoding() {
let manager = Arc::new(VectorManager::new(VectorManagerOptions {
is_enabled: false,
..Default::default()
}));
let sess = RespServerSessionVectors::new(manager);
for r in [
sess.network_vadd(&[b"k", b"FP32", &[0; 4], b"e"]),
sess.network_vsim(&[b"k", b"ELE", b"e"]),
sess.network_vcard(&[b"k"]),
sess.network_vrem(&[b"k", b"e"]),
] {
assert!(matches!(r, VectorReply::Error(_)), "未启用应拒绝");
}
assert!(sess.abort_vector_set_wrong_type(b"somekey").is_none());
let reply = VectorReply::Array(vec![
VectorReply::Bulk(Some(b"a".to_vec())),
VectorReply::Double(2.0),
VectorReply::Boolean(true),
]);
let mut r2 = Vec::new();
reply.encode_resp2(&mut r2);
assert_eq!(s(&r2), "*3\r\n$1\r\na\r\n$1\r\n2\r\n$1\r\n1\r\n");
let mut r3 = Vec::new();
reply.encode_resp3(&mut r3);
assert_eq!(s(&r3), "*3\r\n$1\r\na\r\n,2\r\n#t\r\n");
let map = VectorReply::Map(vec![(
VectorReply::Bulk(Some(b"id".to_vec())),
VectorReply::Double(1.5),
)]);
let mut m2 = Vec::new();
map.encode_resp2(&mut m2);
assert_eq!(s(&m2), "*2\r\n$2\r\nid\r\n$3\r\n1.5\r\n");
let mut m3 = Vec::new();
map.encode_resp3(&mut m3);
assert_eq!(s(&m3), "%1\r\n$2\r\nid\r\n,1.5\r\n");
let mut na = Vec::new();
VectorReply::NullArray.encode_resp2(&mut na);
assert_eq!(s(&na), "*-1\r\n");
}
#[test]
fn result_writers_honor_bitmap_and_count() {
let ids = vec![b"a".as_slice(), b"b".as_slice(), b"c".as_slice()];
let distances = [0.1f32, 0.2, 0.3];
let bitmap = [0b101u8];
let out =
RespServerSessionVectors::write_resp2_result(10, &ids, &distances, &bitmap, None, false);
let mut encoded = Vec::new();
out.encode_resp2(&mut encoded);
assert_eq!(s(&encoded), "*2\r\n$1\r\na\r\n$1\r\nc\r\n");
let out =
RespServerSessionVectors::write_resp3_result(10, &ids, &distances, &bitmap, None, false);
let mut encoded = Vec::new();
out.encode_resp3(&mut encoded);
assert_eq!(s(&encoded), "*2\r\n$1\r\na\r\n$1\r\nc\r\n");
let attrs = vec![b"x".as_slice(), b"".as_slice(), b"z".as_slice()];
let out = RespServerSessionVectors::write_resp3_result(
10,
&ids,
&distances,
&bitmap,
Some(&attrs),
true,
);
let mut encoded = Vec::new();
out.encode_resp3(&mut encoded);
let text = s(&encoded);
assert!(text.starts_with("%2\r\n"), "{text}");
assert!(
text.contains("$1\r\na\r\n*2\r\n,0.10000000149011612\r\n$1\r\nx\r\n"),
"{text}"
);
assert!(
text.contains("$1\r\nc\r\n*2\r\n,0.30000001192092896\r\n$1\r\nz\r\n"),
"{text}"
);
let attrs_one = vec![b"".as_slice()];
let out = RespServerSessionVectors::write_resp3_result(
10,
&[b"e".as_slice()],
&[0.5],
&[],
Some(&attrs_one),
true,
);
let mut encoded = Vec::new();
out.encode_resp3(&mut encoded);
assert_eq!(s(&encoded), "%1\r\n$1\r\ne\r\n*2\r\n,0.5\r\n$-1\r\n");
let out = RespServerSessionVectors::write_resp2_result(1, &ids, &distances, &[], None, false);
let mut encoded = Vec::new();
out.encode_resp2(&mut encoded);
assert_eq!(s(&encoded), "*1\r\n$1\r\na\r\n");
}
fn f32_bytes(vals: &[f32]) -> Vec<u8> {
vals.iter().flat_map(|v| v.to_le_bytes()).collect()
}
fn seed_index(sess: &RespServerSessionVectors, key: &[u8], dims: u32) -> [u8; 56] {
let context = sess.manager.next_vector_set_context(0).unwrap();
sess.manager.service.create_index(
context,
dims,
0,
VectorQuantType::NoQuant,
64,
8,
VectorDistanceMetricType::L2,
);
let index = Index {
context,
index_ptr: 1,
dimensions: dims,
reduce_dims: 0,
num_links: 8,
build_exploration_factor: 64,
quant_type: VectorQuantType::NoQuant,
distance_metric: VectorDistanceMetricType::L2,
..Index::default()
};
let bytes = index.to_bytes();
sess.manager.write_stored_index(key, &bytes);
bytes
}
#[test]
fn id_format_is_length_prefixed() {
assert_eq!(VectorIdFormat::I32LengthPrefixed as i32, 1);
}
}