use crate::db_type::DB_TYPE_BINARY_DOUBLE;
use crate::db_type::DB_TYPE_BINARY_FLOAT;
use crate::db_type::DbType;
use crate::error::Error;
use crate::read_buffer::FromBuf;
use crate::read_buffer::FromBufFallible;
use crate::read_buffer::ReadBuffer;
use crate::write_buffer::ToBuf;
use crate::write_buffer::WriteBuffer;
const VECTOR_MAGIC_BYTE: u8 = 0xDB;
const VECTOR_VERSION_BASE: u8 = 0;
const VECTOR_VERSION_WITH_BINARY: u8 = 1;
const VECTOR_VERSION_WITH_SPARSE: u8 = 2;
const VECTOR_FLAG_NORM: u16 = 0x0002;
const VECTOR_FLAG_NORM_RESERVED: u16 = 0x0010;
const VECTOR_FLAG_SPARSE: u16 = 0x0020;
const VECTOR_FORMAT_BINARY: u8 = 0x05;
const VECTOR_FORMAT_FLOAT32: u8 = 0x02;
const VECTOR_FORMAT_FLOAT64: u8 = 0x03;
const VECTOR_FORMAT_INT8: u8 = 0x04;
#[derive(Debug, Clone)]
pub enum VectorData {
Float32(Vec<f32>),
Float64(Vec<f64>),
Int8(Vec<i8>),
Binary(Vec<u8>),
}
impl VectorData {
fn decode(
buf: &mut ReadBuffer,
num_elements: usize,
vector_format: u8,
) -> Result<Self, Error> {
match vector_format {
VECTOR_FORMAT_FLOAT32 => {
let mut values = Vec::with_capacity(num_elements);
for _ in 0..num_elements {
values.push(f32::from_buf(buf.read_bytes(4)?));
}
Ok(Self::Float32(values))
}
VECTOR_FORMAT_FLOAT64 => {
let mut values = Vec::with_capacity(num_elements);
for _ in 0..num_elements {
values.push(f64::from_buf(buf.read_bytes(8)?));
}
Ok(Self::Float64(values))
}
VECTOR_FORMAT_INT8 => {
let mut values = Vec::with_capacity(num_elements);
for _ in 0..num_elements {
values.push(buf.read_i8()?);
}
Ok(Self::Int8(values))
}
VECTOR_FORMAT_BINARY => {
let byte_count = num_elements / 8;
let mut values = Vec::with_capacity(byte_count);
for _ in 0..byte_count {
values.push(buf.read_u8()?);
}
Ok(Self::Binary(values))
}
_ => Err(Error::unsupported_vector_format(vector_format)),
}
}
fn encode(&self, buf: &mut WriteBuffer) {
match self {
Self::Int8(values) => {
for value in values {
buf.write_u8(*value as u8);
}
}
Self::Binary(values) => {
for value in values {
buf.write_u8(*value);
}
}
Self::Float32(values) => {
for value in values {
value.to_buf(buf, &DB_TYPE_BINARY_FLOAT, false);
}
}
Self::Float64(values) => {
for value in values {
value.to_buf(buf, &DB_TYPE_BINARY_DOUBLE, false);
}
}
}
}
fn format(&self) -> u8 {
match self {
Self::Float32(_) => VECTOR_FORMAT_FLOAT32,
Self::Float64(_) => VECTOR_FORMAT_FLOAT64,
Self::Int8(_) => VECTOR_FORMAT_INT8,
Self::Binary(_) => VECTOR_FORMAT_BINARY,
}
}
fn num_dimensions(&self) -> usize {
match self {
VectorData::Float32(v) => v.len(),
VectorData::Float64(v) => v.len(),
VectorData::Int8(v) => v.len(),
VectorData::Binary(v) => v.len() * 8,
}
}
}
#[derive(Debug, Clone)]
pub struct SparseVector {
num_dimensions: usize,
indices: Vec<usize>,
values: VectorData,
}
impl SparseVector {
fn decode(
buf: &mut ReadBuffer,
num_dimensions: usize,
vector_format: u8,
) -> Result<Self, Error> {
let num_sparse_elements = buf.read_u16be()? as usize;
let mut indices = Vec::with_capacity(num_sparse_elements);
for _ in 0..num_sparse_elements {
indices.push(buf.read_u32be()? as usize);
}
let values =
VectorData::decode(buf, num_sparse_elements, vector_format)?;
Ok(Self {
num_dimensions,
indices,
values,
})
}
fn encode(&self, buf: &mut WriteBuffer) {
buf.write_u16be(self.indices.len().try_into().unwrap());
for ix in &self.indices {
buf.write_u32be((*ix).try_into().unwrap());
}
self.values.encode(buf);
}
pub fn new(
num_dimensions: usize,
indices: Vec<usize>,
values: VectorData,
) -> Self {
Self {
num_dimensions,
indices,
values,
}
}
pub fn num_dimensions(&self) -> usize {
self.num_dimensions
}
pub fn indices(&self) -> &[usize] {
&self.indices
}
pub fn values(&self) -> &VectorData {
&self.values
}
}
#[derive(Debug, Clone)]
pub enum Vector {
Dense(VectorData),
Sparse(SparseVector),
}
impl Vector {
fn decode(buf: &mut ReadBuffer) -> Result<Self, Error> {
let magic_byte = buf.read_u8()?;
if magic_byte != VECTOR_MAGIC_BYTE {
return Err(Error::invalid_encoded_vector());
}
let version = buf.read_u8()?;
if version > VECTOR_VERSION_WITH_SPARSE {
return Err(Error::unsupported_vector_version(version));
}
let flags = buf.read_u16be()?;
let vector_format = buf.read_u8()?;
let num_elements = buf.read_u32be()? as usize;
if (flags & VECTOR_FLAG_NORM_RESERVED) != 0
|| (flags & VECTOR_FLAG_NORM) != 0
{
buf.read_bytes(8)?;
}
if flags & VECTOR_FLAG_SPARSE != 0 {
let sparse =
SparseVector::decode(buf, num_elements, vector_format)?;
Ok(Self::Sparse(sparse))
} else {
let values = VectorData::decode(buf, num_elements, vector_format)?;
Ok(Self::Dense(values))
}
}
fn flags(&self) -> u16 {
let base_flags = VECTOR_FLAG_NORM_RESERVED | VECTOR_FLAG_NORM;
match self {
Self::Dense(_) => base_flags,
Self::Sparse(_) => base_flags | VECTOR_FLAG_SPARSE,
}
}
fn format(&self) -> u8 {
match self {
Self::Dense(data) => data.format(),
Self::Sparse(sparse) => sparse.values.format(),
}
}
fn num_dimensions(&self) -> usize {
match self {
Self::Dense(data) => data.num_dimensions(),
Self::Sparse(sparse) => sparse.num_dimensions,
}
}
fn version(&self) -> u8 {
match self {
Self::Dense(data) => match data.format() {
VECTOR_FORMAT_BINARY => VECTOR_VERSION_WITH_BINARY,
_ => VECTOR_VERSION_BASE,
},
Self::Sparse(_) => VECTOR_VERSION_WITH_SPARSE,
}
}
pub(crate) fn encode(&self, buf: &mut WriteBuffer) {
buf.write_u8(VECTOR_MAGIC_BYTE);
buf.write_u8(self.version());
buf.write_u16be(self.flags());
buf.write_u8(self.format());
buf.write_u32be(self.num_dimensions().try_into().unwrap());
buf.write_bytes(&[0u8; 8]);
match self {
Vector::Dense(data) => {
data.encode(buf);
}
Vector::Sparse(sparse) => {
sparse.encode(buf);
}
}
}
}
impl FromBufFallible for Vector {
fn from_buf_fallible(buf: &mut ReadBuffer) -> Result<Self, Error> {
Vector::decode(buf)
}
}
impl ToBuf for Vector {
fn to_buf(
&self,
buf: &mut WriteBuffer,
_db_type: &'static DbType,
_write_length: bool,
) {
let mut encode_buf = WriteBuffer::new();
self.encode(&mut encode_buf);
buf.write_qlocator(&encode_buf);
}
}