use std::sync::atomic::{AtomicU32, Ordering};
use super::{disk_ann_service::term, vector_manager::VectorManager, vector_types::VectorQuantType};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum ReadCopyTo {
#[default]
ReadCache = 0,
MainLog = 1,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub struct VectorReadGeometry {
pub full_vector_io_size: usize,
pub neighbor_list_io_size: usize,
pub quantized_vector_io_size: usize,
}
const LINK_BYTES: usize = 4;
const F32_SIZE: usize = 4;
const VECTOR_RECORD_READ_OVERHEAD_BYTES: usize = 64;
static ACTIVE_FULL_VECTOR_IO: AtomicU32 = AtomicU32::new(0);
static ACTIVE_NEIGHBOR_LIST_IO: AtomicU32 = AtomicU32::new(0);
static ACTIVE_QUANTIZED_IO: AtomicU32 = AtomicU32::new(0);
pub fn set_active_read_geometry(
dims: u32,
num_links: u32,
quant: VectorQuantType,
reduce_dims: u32,
) {
let effective_dims = if reduce_dims != 0 {
reduce_dims.min(dims)
} else {
dims
};
let element_size = match quant {
VectorQuantType::XNoQuant_U8
| VectorQuantType::XBin_U8
| VectorQuantType::XNoQuant_I8
| VectorQuantType::XBin_I8 => 1,
_ => F32_SIZE,
};
let full = effective_dims as usize * element_size + VECTOR_RECORD_READ_OVERHEAD_BYTES;
let neighbor_list = num_links as usize * 2 * LINK_BYTES + VECTOR_RECORD_READ_OVERHEAD_BYTES;
let quantized = match quant {
VectorQuantType::Q8 => dims as usize + VECTOR_RECORD_READ_OVERHEAD_BYTES,
_ => full,
};
ACTIVE_FULL_VECTOR_IO.store(clamp_io(full), Ordering::Release);
ACTIVE_NEIGHBOR_LIST_IO.store(clamp_io(neighbor_list), Ordering::Release);
ACTIVE_QUANTIZED_IO.store(clamp_io(quantized), Ordering::Release);
}
fn clamp_io(size: usize) -> u32 {
size.min(u32::MAX as usize) as u32
}
pub fn active_read_geometry() -> VectorReadGeometry {
VectorReadGeometry {
full_vector_io_size: ACTIVE_FULL_VECTOR_IO.load(Ordering::Acquire) as usize,
neighbor_list_io_size: ACTIVE_NEIGHBOR_LIST_IO.load(Ordering::Acquire) as usize,
quantized_vector_io_size: ACTIVE_QUANTIZED_IO.load(Ordering::Acquire) as usize,
}
}
#[derive(Debug, Clone, Default)]
pub struct VectorReadBatch<'a> {
parameters: &'a [u8],
keys: Vec<&'a [u8]>,
current_index: usize,
}
impl<'a> VectorReadBatch<'a> {
pub fn new(parameters: &'a [u8]) -> Self {
let mut batch = Self {
parameters,
keys: Vec::new(),
current_index: 0,
};
batch.parse_all();
batch
}
fn parse_all(&mut self) {
self.keys.clear();
let mut rest = self.parameters;
while rest.len() >= 4 {
let len = i32::from_le_bytes(rest[..4].try_into().unwrap_or([0; 4]));
if len < 0 {
break;
}
let total = 4 + len as usize;
if rest.len() < total {
break;
}
self.keys.push(&rest[4..total]);
rest = &rest[total..];
}
}
pub fn count(&self) -> usize {
self.keys.len()
}
pub fn advance_to(&mut self, i: usize) -> bool {
if i >= self.keys.len() {
return false;
}
self.current_index = i;
true
}
pub fn get_key(&self) -> Option<&'a [u8]> {
self.keys.get(self.current_index).copied()
}
pub fn key_at(&self, i: usize) -> Option<&'a [u8]> {
self.keys.get(i).copied()
}
pub fn get_input(&self) -> Option<&'a [u8]> {
self.get_key()
}
pub fn get_output(&self) -> usize {
active_read_geometry().full_vector_io_size
}
pub fn set_output(&self, output: &mut [u8], record: &[u8]) {
let n = output.len().min(record.len());
output[..n].copy_from_slice(&record[..n]);
}
}
impl VectorManager {
pub fn make_vector_element_key(namespace_bytes: &[u8], key_data: &[u8]) -> Vec<u8> {
let mut out = Vec::with_capacity(namespace_bytes.len() + key_data.len());
out.extend_from_slice(namespace_bytes);
out.extend_from_slice(key_data);
out
}
pub fn read_callback(&self, context_term: u64, key: &[u8]) -> Option<Vec<u8>> {
let context = context_term & !term_mask();
match context_term & term_mask() {
t if t == term::ATTRIBUTES => self.service.get_attribute(context, key),
t if t == term::FULL_VECTOR => self.service.get_full_vector(context, key),
t if t == term::INTERNAL_ID_MAP => self
.service
.internal_id_of(context, key)
.map(|id| id.to_le_bytes().to_vec()),
_ => None,
}
}
pub fn write_callback(&self, context_term: u64, key: &[u8], value: &[u8]) -> bool {
if context_term & term_mask() == term::ATTRIBUTES {
return self
.service
.set_attribute(context_term & !term_mask(), key, value);
}
false
}
pub fn delete_callback(&self, context_term: u64, key: &[u8]) -> bool {
if context_term & term_mask() == term::ATTRIBUTES {
return self.service.remove(context_term & !term_mask(), key);
}
false
}
pub fn read_modify_write_callback(
&self,
context_term: u64,
key: &[u8],
upsert: impl FnOnce(Option<&[u8]>) -> Vec<u8>,
) -> Option<Vec<u8>> {
let current = self.read_callback(context_term, key);
let next = upsert(current.as_deref());
let _ = self.write_callback(context_term, key, &next);
Some(next)
}
pub fn filter_callback(&self, context_term: u64, internal_id: u32) -> bool {
let context = context_term & !term_mask();
if context_term & term_mask() != term::ATTRIBUTES {
return false;
}
self.service.check_internal_id_valid(context, internal_id)
}
pub fn read_size_unknown(&self, context_term: u64, key: &[u8]) -> Option<Vec<u8>> {
self.read_callback(context_term, key)
}
pub fn slow_path(&self, context_term: u64, batch: &VectorReadBatch<'_>) -> Vec<Option<Vec<u8>>> {
(0..batch.count())
.map(|i| {
let key = batch.key_at(i)?;
self.read_size_unknown(context_term, key)
})
.collect()
}
}
fn term_mask() -> u64 {
0b111
}
#[cfg(test)]
mod tests {
use super::{
super::{
vector_manager::{VectorManager, VectorManagerOptions},
vector_types::{VectorDistanceMetricType, VectorQuantType},
},
*,
};
fn manager() -> VectorManager {
VectorManager::new(VectorManagerOptions {
is_enabled: true,
..Default::default()
})
}
#[test]
fn geometry_sizes_follow_config() {
set_active_read_geometry(8, 4, VectorQuantType::NoQuant, 0);
let g = active_read_geometry();
assert_eq!(g.full_vector_io_size, 8 * 4 + 64);
assert_eq!(g.neighbor_list_io_size, 4 * 2 * 4 + 64);
set_active_read_geometry(8, 4, VectorQuantType::Q8, 0);
let g = active_read_geometry();
assert_eq!(g.quantized_vector_io_size, 8 + 64);
set_active_read_geometry(16, 4, VectorQuantType::NoQuant, 4);
assert_eq!(active_read_geometry().full_vector_io_size, 4 * 4 + 64);
}
#[test]
fn read_batch_parses_length_prefixed_stream() {
let mut stream: Vec<u8> = Vec::new();
for key in [&b"alpha"[..], &b"b"[..], &b"gamma"[..]] {
stream.extend_from_slice(&(key.len() as i32).to_le_bytes());
stream.extend_from_slice(key);
}
let mut batch = VectorReadBatch::new(&stream);
assert_eq!(batch.count(), 3);
assert!(batch.advance_to(0));
assert_eq!(batch.get_key(), Some(b"alpha".as_slice()));
assert_eq!(batch.get_input(), Some(b"alpha".as_slice()));
assert!(batch.advance_to(2));
assert_eq!(batch.get_key(), Some(b"gamma".as_slice()));
assert!(!batch.advance_to(9));
assert_eq!(batch.get_key(), Some(b"gamma".as_slice()));
set_active_read_geometry(4, 2, VectorQuantType::NoQuant, 0);
assert_eq!(batch.get_output(), 4 * 4 + 64);
let mut out = [0u8; 3];
batch.set_output(&mut out, b"abcdef");
assert_eq!(&out, b"abc");
let truncated = VectorReadBatch::new(&stream[..stream.len() - 1]);
assert_eq!(truncated.count(), 2);
}
#[test]
fn element_key_composition() {
let key = VectorManager::make_vector_element_key(&[4, 0, 0, 0], b"elem");
assert_eq!(key, vec![4, 0, 0, 0, b'e', b'l', b'e', b'm']);
assert_eq!(
VectorManager::make_vector_element_key(&[], b"x"),
vec![b'x']
);
}
#[test]
fn term_dispatch_read_write_delete() {
let manager = manager();
manager.service.create_index(
16,
1,
0,
VectorQuantType::NoQuant,
16,
2,
VectorDistanceMetricType::L2,
);
manager
.service
.insert(16, b"e", &1.0f32.to_le_bytes(), b"{\"a\":1}");
let attr_term = 16 | term::ATTRIBUTES;
assert_eq!(
manager.read_callback(attr_term, b"e"),
Some(b"{\"a\":1}".to_vec())
);
assert!(manager.write_callback(attr_term, b"e", b"{}"));
assert_eq!(manager.read_callback(attr_term, b"e"), Some(b"{}".to_vec()));
let full_term = 16 | term::FULL_VECTOR;
assert!(manager.read_callback(full_term, b"e").is_some());
let next = manager.read_modify_write_callback(attr_term, b"e", |prev| {
let mut v = prev.unwrap_or_default().to_vec();
v.extend_from_slice(b"#");
v
});
assert_eq!(next, Some(b"{}#".to_vec()));
assert!(manager.delete_callback(attr_term, b"e"));
assert_eq!(manager.read_callback(attr_term, b"e"), None);
assert!(!manager.write_callback(full_term, b"e", b"x"));
assert_eq!(manager.read_size_unknown(attr_term, b"e"), None);
}
#[test]
fn filter_callback_requires_attributes_term() {
let manager = manager();
manager.service.create_index(
8,
1,
0,
VectorQuantType::NoQuant,
8,
2,
VectorDistanceMetricType::L2,
);
manager.service.insert(8, b"x", &5f32.to_le_bytes(), b"{}");
assert!(manager.filter_callback(8 | term::ATTRIBUTES, 0));
assert!(!manager.filter_callback(8 | term::ATTRIBUTES, 99));
assert!(!manager.filter_callback(8 | term::FULL_VECTOR, 0));
}
#[test]
fn slow_path_reads_per_key_with_placeholders() {
let manager = manager();
manager.service.create_index(
24,
1,
0,
VectorQuantType::NoQuant,
16,
2,
VectorDistanceMetricType::L2,
);
manager
.service
.insert(24, b"hit", &1f32.to_le_bytes(), b"{}");
let mut stream: Vec<u8> = Vec::new();
for key in [&b"miss"[..], &b"hit"[..]] {
stream.extend_from_slice(&(key.len() as i32).to_le_bytes());
stream.extend_from_slice(key);
}
let batch = VectorReadBatch::new(&stream);
let results = manager.slow_path(24 | term::ATTRIBUTES, &batch);
assert_eq!(results.len(), 2);
assert_eq!(results[0], None);
assert_eq!(results[1], Some(b"{}".to_vec()));
}
}