Skip to main content

kevy_vector/
dist.rs

1//! Distance metrics + the wire vector format (RFC D1/D3).
2
3/// Distance metric. Scores are "smaller = closer" for every variant
4/// (cosine → `1 - cos`, ip → `-dot`), so one ascending merge works.
5#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
6pub enum Distance {
7    /// Cosine distance (vectors pre-normalized at insert).
8    #[default]
9    Cosine,
10    /// Squared euclidean.
11    L2,
12    /// Negative inner product.
13    Ip,
14}
15
16impl Distance {
17    /// Tag for sidecar round-trip.
18    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    /// Parse a tag (ASCII case-insensitive).
27    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    /// Distance between two prepared vectors (see [`prepare`]).
40    #[inline]
41    pub(crate) fn eval(self, a: &[f32], b: &[f32]) -> f32 {
42        match self {
43            // prepared cosine vectors are unit length → 1 - dot
44            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    /// Normalize a vector into its stored/query form (cosine only).
51    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
68/// Decode a wire vector: raw f32 LE bytes (`len == dim*4`), or the
69/// debug form `csv:1.0,2.5,…`. `None` on any mismatch.
70pub 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}