Skip to main content

radixdb_core/
vector.rs

1// Copyright 2026 RadixDB Contributors
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8
9//! Allocation-free vector distance primitives shared by storage and SQL.
10
11use crate::{Error, Result};
12
13#[inline]
14pub fn validate_vector_bytes(data: &[u8]) -> Result<()> {
15    if !data.len().is_multiple_of(4) {
16        return Err(Error::invalid_argument(format!(
17            "malformed VECTOR payload: {} bytes is not aligned to f32",
18            data.len()
19        )));
20    }
21    Ok(())
22}
23
24#[inline]
25fn validate_equal_bytes(a: &[u8], b: &[u8]) -> Result<()> {
26    validate_vector_bytes(a)?;
27    validate_vector_bytes(b)?;
28    if a.len() != b.len() {
29        return Err(Error::invalid_argument(format!(
30            "vector dimension mismatch ({} vs {})",
31            a.len() / 4,
32            b.len() / 4
33        )));
34    }
35    Ok(())
36}
37
38#[inline(always)]
39fn read_f32(data: &[u8], index: usize) -> f32 {
40    let offset = index * 4;
41    f32::from_le_bytes([
42        data[offset],
43        data[offset + 1],
44        data[offset + 2],
45        data[offset + 3],
46    ])
47}
48
49/// L2 (Euclidean) distance on raw little-endian f32 byte slices.
50#[inline]
51pub fn l2_distance_bytes(a: &[u8], b: &[u8]) -> Result<f64> {
52    validate_equal_bytes(a, b)?;
53    let len = a.len() / 4;
54    let mut sum = 0.0f64;
55    let mut index = 0;
56    while index + 4 <= len {
57        let d0 = (read_f32(a, index) - read_f32(b, index)) as f64;
58        let d1 = (read_f32(a, index + 1) - read_f32(b, index + 1)) as f64;
59        let d2 = (read_f32(a, index + 2) - read_f32(b, index + 2)) as f64;
60        let d3 = (read_f32(a, index + 3) - read_f32(b, index + 3)) as f64;
61        sum += d0 * d0 + d1 * d1 + d2 * d2 + d3 * d3;
62        index += 4;
63    }
64    while index < len {
65        let distance = (read_f32(a, index) - read_f32(b, index)) as f64;
66        sum += distance * distance;
67        index += 1;
68    }
69    Ok(sum.sqrt())
70}
71
72/// Cosine distance (`1 - cosine_similarity`) on raw LE f32 bytes.
73#[inline]
74pub fn cosine_distance_bytes(a: &[u8], b: &[u8]) -> Result<f64> {
75    validate_equal_bytes(a, b)?;
76    let len = a.len() / 4;
77    let mut dot = 0.0f64;
78    let mut norm_a = 0.0f64;
79    let mut norm_b = 0.0f64;
80    for index in 0..len {
81        let ai = read_f32(a, index) as f64;
82        let bi = read_f32(b, index) as f64;
83        dot += ai * bi;
84        norm_a += ai * ai;
85        norm_b += bi * bi;
86    }
87    let denominator = norm_a.sqrt() * norm_b.sqrt();
88    if denominator == 0.0 {
89        Ok(1.0)
90    } else {
91        Ok((1.0 - (dot / denominator)).max(0.0))
92    }
93}
94
95/// Negative inner-product distance on raw LE f32 bytes.
96#[inline]
97pub fn ip_distance_bytes(a: &[u8], b: &[u8]) -> Result<f64> {
98    validate_equal_bytes(a, b)?;
99    let len = a.len() / 4;
100    let mut dot = 0.0f64;
101    for index in 0..len {
102        dot += (read_f32(a, index) as f64) * (read_f32(b, index) as f64);
103    }
104    Ok(-dot)
105}
106
107/// L2 (Euclidean) distance on decoded f32 slices.
108#[inline]
109pub fn l2_distance(a: &[f32], b: &[f32]) -> Result<f64> {
110    if a.len() != b.len() {
111        return Err(Error::invalid_argument(format!(
112            "vector dimension mismatch ({} vs {})",
113            a.len(),
114            b.len()
115        )));
116    }
117    let mut sum = 0.0f64;
118    let len = a.len();
119    let mut index = 0;
120    while index + 4 <= len {
121        let d0 = (a[index] - b[index]) as f64;
122        let d1 = (a[index + 1] - b[index + 1]) as f64;
123        let d2 = (a[index + 2] - b[index + 2]) as f64;
124        let d3 = (a[index + 3] - b[index + 3]) as f64;
125        sum += d0 * d0 + d1 * d1 + d2 * d2 + d3 * d3;
126        index += 4;
127    }
128    while index < len {
129        let distance = (a[index] - b[index]) as f64;
130        sum += distance * distance;
131        index += 1;
132    }
133    Ok(sum.sqrt())
134}
135
136/// Cosine distance on decoded f32 slices.
137#[inline]
138pub fn cosine_distance(a: &[f32], b: &[f32]) -> Result<f64> {
139    if a.len() != b.len() {
140        return Err(Error::invalid_argument(format!(
141            "vector dimension mismatch ({} vs {})",
142            a.len(),
143            b.len()
144        )));
145    }
146    let mut dot = 0.0f64;
147    let mut norm_a = 0.0f64;
148    let mut norm_b = 0.0f64;
149    for index in 0..a.len() {
150        let ai = a[index] as f64;
151        let bi = b[index] as f64;
152        dot += ai * bi;
153        norm_a += ai * ai;
154        norm_b += bi * bi;
155    }
156    let denominator = norm_a.sqrt() * norm_b.sqrt();
157    if denominator == 0.0 {
158        Ok(1.0)
159    } else {
160        Ok((1.0 - (dot / denominator)).max(0.0))
161    }
162}
163
164#[cfg(test)]
165mod tests {
166    use super::*;
167
168    fn encode(values: &[f32]) -> Vec<u8> {
169        values
170            .iter()
171            .flat_map(|value| value.to_le_bytes())
172            .collect()
173    }
174
175    #[test]
176    fn byte_and_decoded_distances_are_identical() {
177        let a = [1.0, -2.0, 3.5, 4.0, 9.0];
178        let b = [-1.0, 2.0, 1.5, 8.0, 3.0];
179        let encoded_a = encode(&a);
180        let encoded_b = encode(&b);
181
182        assert_eq!(
183            l2_distance_bytes(&encoded_a, &encoded_b).unwrap(),
184            l2_distance(&a, &b).unwrap()
185        );
186        assert_eq!(
187            cosine_distance_bytes(&encoded_a, &encoded_b).unwrap(),
188            cosine_distance(&a, &b).unwrap()
189        );
190    }
191
192    #[test]
193    fn malformed_and_mismatched_payloads_are_rejected() {
194        assert!(l2_distance_bytes(&[0, 0, 0], &[0, 0, 0]).is_err());
195        assert!(cosine_distance_bytes(&encode(&[1.0]), &encode(&[1.0, 2.0])).is_err());
196        assert!(ip_distance_bytes(&encode(&[1.0]), &encode(&[1.0, 2.0])).is_err());
197    }
198
199    #[test]
200    fn zero_vector_cosine_contract_is_stable() {
201        assert_eq!(cosine_distance(&[0.0, 0.0], &[1.0, 2.0]).unwrap(), 1.0);
202        assert_eq!(
203            cosine_distance_bytes(&encode(&[0.0, 0.0]), &encode(&[1.0, 2.0])).unwrap(),
204            1.0
205        );
206    }
207}