use crate::{get_u8, get_u16, get_u32, put_u8, put_u16, put_u32};
use yo_common::{Code, Error, Result};
pub const VECTOR_HEADER_LEN: usize = 8;
pub const MAX_DIM: usize = 65_536;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u8)]
pub enum Element {
F32 = 0,
}
impl Element {
#[must_use]
pub const fn as_u8(self) -> u8 {
self as u8
}
#[must_use]
pub const fn from_u8(b: u8) -> Option<Element> {
match b {
0 => Some(Element::F32),
_ => None,
}
}
#[must_use]
pub const fn width(self) -> usize {
match self {
Element::F32 => 4,
}
}
}
pub fn vector_len(dim: usize, of: Element) -> Result<usize> {
if dim == 0 || dim > MAX_DIM {
return Err(Error::new(Code::Invalid, "dimension out of range")
.with_detail(format!("dim={dim} max={MAX_DIM}")));
}
Ok(VECTOR_HEADER_LEN + dim * of.width())
}
#[derive(Debug, Clone, Copy)]
pub struct VectorBody<'a> {
dim: usize,
element: Element,
values: &'a [u8],
}
impl<'a> VectorBody<'a> {
pub fn encode(values: &[f32], into: &mut [u8]) -> Result<usize> {
let need = vector_len(values.len(), Element::F32)?;
if into.len() < need {
return Err(
Error::new(Code::Invalid, "buffer is shorter than the vector")
.with_detail(format!("have={} need={need}", into.len())),
);
}
if let Some(at) = values.iter().position(|v| !v.is_finite()) {
return Err(Error::new(Code::Invalid, "a coordinate is not a number")
.with_detail(format!("at={at} value={}", values[at])));
}
put_u32(into, 0, values.len() as u32);
put_u8(into, 4, Element::F32.as_u8());
put_u8(into, 5, 0);
put_u16(into, 6, 0);
for (i, v) in values.iter().enumerate() {
let at = VECTOR_HEADER_LEN + i * 4;
into[at..at + 4].copy_from_slice(&v.to_le_bytes());
}
Ok(need)
}
pub fn decode(bytes: &'a [u8]) -> Result<VectorBody<'a>> {
if bytes.len() < VECTOR_HEADER_LEN {
return Err(Error::new(Code::Corrupt, "shorter than a vector header")
.with_detail(format!("len={}", bytes.len())));
}
let dim = get_u32(bytes, 0) as usize;
let raw = get_u8(bytes, 4);
let Some(element) = Element::from_u8(raw) else {
return Err(Error::new(Code::Corrupt, "unknown vector element type")
.with_detail(format!("element={raw}")));
};
let flags = get_u8(bytes, 5);
let reserved = get_u16(bytes, 6);
if flags != 0 || reserved != 0 {
return Err(
Error::new(Code::Corrupt, "reserved vector header bytes are set")
.with_detail(format!("flags={flags:#04x} reserved={reserved:#06x}")),
);
}
if dim == 0 || dim > MAX_DIM {
return Err(Error::new(Code::Corrupt, "vector dimension out of range")
.with_detail(format!("dim={dim} max={MAX_DIM}")));
}
let need = VECTOR_HEADER_LEN + dim * element.width();
if bytes.len() < need {
return Err(
Error::new(Code::Corrupt, "vector record is shorter than its dimension")
.with_detail(format!("len={} need={need} dim={dim}", bytes.len())),
);
}
Ok(VectorBody {
dim,
element,
values: &bytes[VECTOR_HEADER_LEN..need],
})
}
#[must_use]
pub const fn dim(&self) -> usize {
self.dim
}
#[must_use]
pub const fn element(&self) -> Element {
self.element
}
pub fn read_into(&self, out: &mut [f32]) -> Result<()> {
if out.len() != self.dim {
return Err(
Error::new(Code::Invalid, "buffer is not the vector's length")
.with_detail(format!("have={} dim={}", out.len(), self.dim)),
);
}
match self.element {
Element::F32 => {
for (i, slot) in out.iter_mut().enumerate() {
let at = i * 4;
let b = &self.values[at..at + 4];
*slot = f32::from_le_bytes([b[0], b[1], b[2], b[3]]);
}
}
}
Ok(())
}
#[must_use]
pub fn at(&self, i: usize) -> Option<f32> {
if i >= self.dim {
return None;
}
match self.element {
Element::F32 => {
let b = self.values.get(i * 4..i * 4 + 4)?;
Some(f32::from_le_bytes([b[0], b[1], b[2], b[3]]))
}
}
}
#[must_use]
pub const fn bytes(&self) -> &'a [u8] {
self.values
}
}
#[cfg(test)]
mod tests {
use super::*;
fn round_trip(values: &[f32]) -> Vec<f32> {
let mut buf = vec![0u8; vector_len(values.len(), Element::F32).unwrap()];
let wrote = VectorBody::encode(values, &mut buf).unwrap();
assert_eq!(
wrote,
buf.len(),
"encode wrote a different length than it asked for"
);
let body = VectorBody::decode(&buf).unwrap();
assert_eq!(body.dim(), values.len());
assert_eq!(body.element(), Element::F32);
let mut out = vec![0f32; body.dim()];
body.read_into(&mut out).unwrap();
out
}
#[test]
fn a_vector_comes_back_bit_for_bit() {
let values = [0.0, -0.0, 1.0, -1.0, 1e-38, 3.4e38, 0.1, 2.5];
assert_eq!(round_trip(&values), values);
}
#[test]
fn a_long_vector_is_fine() {
let values: Vec<f32> = (0..1536).map(|i| i as f32 * 0.001).collect();
assert_eq!(round_trip(&values), values);
}
#[test]
fn coordinates_can_be_read_one_at_a_time() {
let values = [3.0f32, 1.0, 4.0, 1.5];
let mut buf = vec![0u8; vector_len(4, Element::F32).unwrap()];
VectorBody::encode(&values, &mut buf).unwrap();
let body = VectorBody::decode(&buf).unwrap();
for (i, want) in values.iter().enumerate() {
assert_eq!(body.at(i), Some(*want));
}
assert_eq!(body.at(4), None, "past the end is not a coordinate");
}
#[test]
fn a_vector_that_is_not_a_vector_is_refused() {
let mut buf = vec![0u8; 64];
assert!(VectorBody::encode(&[], &mut buf).is_err(), "no dimension");
assert!(
VectorBody::encode(&[f32::NAN, 1.0], &mut buf).is_err(),
"a NaN coordinate poisons every distance it takes part in"
);
assert!(VectorBody::encode(&[f32::INFINITY], &mut buf).is_err());
let mut tiny = [0u8; 8];
assert!(
VectorBody::encode(&[1.0, 2.0], &mut tiny).is_err(),
"the header fits and the coordinates do not"
);
}
#[test]
fn a_record_shorter_than_it_claims_is_corrupt() {
let values = [1.0f32, 2.0, 3.0, 4.0];
let mut buf = vec![0u8; vector_len(4, Element::F32).unwrap()];
VectorBody::encode(&values, &mut buf).unwrap();
for len in 0..buf.len() {
assert!(
VectorBody::decode(&buf[..len]).is_err(),
"{len} bytes decoded as a four dimensional vector"
);
}
assert!(VectorBody::decode(&buf).is_ok());
}
#[test]
fn an_element_type_this_version_does_not_know_is_refused() {
let mut buf = vec![0u8; vector_len(2, Element::F32).unwrap()];
VectorBody::encode(&[1.0, 2.0], &mut buf).unwrap();
buf[4] = 1;
let e = VectorBody::decode(&buf).unwrap_err();
assert_eq!(e.code(), Code::Corrupt);
}
#[test]
fn reserved_bytes_have_to_be_zero() {
let values = [1.0f32, 2.0];
for at in [5usize, 6, 7] {
let mut buf = vec![0u8; vector_len(2, Element::F32).unwrap()];
VectorBody::encode(&values, &mut buf).unwrap();
buf[at] = 1;
assert!(
VectorBody::decode(&buf).is_err(),
"byte {at} is reserved and a set bit in it means the writer disagreed with this layout"
);
}
}
#[test]
fn a_dimension_that_could_not_fit_anywhere_is_refused_before_it_is_believed() {
let mut buf = vec![0u8; vector_len(2, Element::F32).unwrap()];
VectorBody::encode(&[1.0, 2.0], &mut buf).unwrap();
put_u32(&mut buf, 0, u32::MAX);
let e = VectorBody::decode(&buf).unwrap_err();
assert_eq!(e.code(), Code::Corrupt);
assert!(vector_len(MAX_DIM + 1, Element::F32).is_err());
}
#[test]
fn reading_into_the_wrong_length_says_so() {
let mut buf = vec![0u8; vector_len(3, Element::F32).unwrap()];
VectorBody::encode(&[1.0, 2.0, 3.0], &mut buf).unwrap();
let body = VectorBody::decode(&buf).unwrap();
assert!(body.read_into(&mut [0.0; 2]).is_err(), "too short");
assert!(body.read_into(&mut [0.0; 4]).is_err(), "too long");
assert!(body.read_into(&mut [0.0; 3]).is_ok());
}
#[test]
fn the_element_byte_maps_both_ways() {
assert_eq!(Element::from_u8(Element::F32.as_u8()), Some(Element::F32));
assert_eq!(Element::from_u8(1), None);
assert_eq!(Element::F32.width(), 4);
}
}