1use 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#[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#[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#[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#[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#[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}