1#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
6pub enum Distance {
7 #[default]
9 Cosine,
10 L2,
12 Ip,
14}
15
16impl Distance {
17 pub fn tag(self) -> &'static str {
19 match self {
20 Distance::Cosine => "cosine",
21 Distance::L2 => "l2",
22 Distance::Ip => "ip",
23 }
24 }
25
26 pub fn parse(raw: &[u8]) -> Option<Distance> {
28 if raw.eq_ignore_ascii_case(b"cosine") {
29 Some(Distance::Cosine)
30 } else if raw.eq_ignore_ascii_case(b"l2") {
31 Some(Distance::L2)
32 } else if raw.eq_ignore_ascii_case(b"ip") {
33 Some(Distance::Ip)
34 } else {
35 None
36 }
37 }
38
39 #[inline]
41 pub(crate) fn eval(self, a: &[f32], b: &[f32]) -> f32 {
42 match self {
43 Distance::Cosine => 1.0 - dot(a, b),
45 Distance::L2 => a.iter().zip(b).map(|(x, y)| (x - y) * (x - y)).sum(),
46 Distance::Ip => -dot(a, b),
47 }
48 }
49
50 pub(crate) fn prepare(self, v: &mut [f32]) {
52 if self == Distance::Cosine {
53 let norm = dot(v, v).sqrt();
54 if norm > 0.0 {
55 for x in v.iter_mut() {
56 *x /= norm;
57 }
58 }
59 }
60 }
61}
62
63#[inline]
64fn dot(a: &[f32], b: &[f32]) -> f32 {
65 a.iter().zip(b).map(|(x, y)| x * y).sum()
66}
67
68pub fn parse_vector(raw: &[u8], dim: usize) -> Option<Vec<f32>> {
71 if let Some(csv) = raw.strip_prefix(b"csv:") {
72 let vals: Option<Vec<f32>> = std::str::from_utf8(csv)
73 .ok()?
74 .split(',')
75 .map(|s| s.trim().parse::<f32>().ok())
76 .collect();
77 let vals = vals?;
78 return (vals.len() == dim && vals.iter().all(|x| x.is_finite())).then_some(vals);
79 }
80 if raw.len() != dim * 4 {
81 return None;
82 }
83 let mut out = Vec::with_capacity(dim);
84 for chunk in raw.chunks_exact(4) {
85 let x = f32::from_le_bytes(chunk.try_into().expect("4 bytes"));
86 if !x.is_finite() {
87 return None;
88 }
89 out.push(x);
90 }
91 Some(out)
92}
93
94#[cfg(test)]
95mod tests {
96 use super::*;
97
98 #[test]
99 fn metrics_smaller_is_closer() {
100 let mut a = vec![1.0, 0.0];
101 let mut b = vec![0.9, 0.1];
102 let mut c = vec![-1.0, 0.0];
103 for v in [&mut a, &mut b, &mut c] {
104 Distance::Cosine.prepare(v);
105 }
106 assert!(Distance::Cosine.eval(&a, &b) < Distance::Cosine.eval(&a, &c));
107 assert!(Distance::L2.eval(&[0.0, 0.0], &[1.0, 1.0]) > Distance::L2.eval(&[0.0, 0.0], &[0.5, 0.5]));
108 assert!(Distance::Ip.eval(&[1.0, 1.0], &[2.0, 2.0]) < Distance::Ip.eval(&[1.0, 1.0], &[0.1, 0.1]));
109 }
110
111 #[test]
112 fn wire_formats() {
113 let mut raw = Vec::new();
114 for x in [1.0f32, -2.5, 3.25] {
115 raw.extend_from_slice(&x.to_le_bytes());
116 }
117 assert_eq!(parse_vector(&raw, 3), Some(vec![1.0, -2.5, 3.25]));
118 assert_eq!(parse_vector(&raw, 4), None, "dim mismatch");
119 assert_eq!(parse_vector(b"csv:1, 2.5, -3", 3), Some(vec![1.0, 2.5, -3.0]));
120 assert_eq!(parse_vector(b"csv:1,x,3", 3), None);
121 let mut nan = Vec::new();
122 for x in [1.0f32, f32::NAN] {
123 nan.extend_from_slice(&x.to_le_bytes());
124 }
125 assert_eq!(parse_vector(&nan, 2), None, "non-finite rejected");
126 assert!(Distance::parse(b"COSINE") == Some(Distance::Cosine));
127 assert!(Distance::parse(b"nope").is_none());
128 }
129}