Skip to main content

khive_types/
vector.rs

1//! Vector similarity and search primitives.
2
3use alloc::{vec, vec::Vec};
4
5mod codec;
6pub use codec::{
7    decode_f32_le, decode_f32_native, encode_f32_le, encode_f32_native, VectorCodecError,
8};
9
10/// Distance metric for vector similarity search.
11///
12/// # Variants
13/// - `Cosine`: `1 - cosine_similarity`. Value in [0, 2] for unit vectors.
14/// - `Dot`: dot product (negated for min-heap; higher dot = lower distance).
15/// - `L2`: Euclidean (L2) distance.
16#[non_exhaustive]
17#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
18#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
19#[cfg_attr(feature = "serde", serde(rename_all = "snake_case"))]
20pub enum DistanceMetric {
21    Cosine,
22    Dot,
23    L2,
24}
25
26impl Default for DistanceMetric {
27    #[inline]
28    fn default() -> Self {
29        Self::Cosine
30    }
31}
32
33/// Generation-based visited-node tracker for greedy search.
34///
35/// Avoids clearing a `Vec<bool>` on every query by incrementing a generation counter.
36#[derive(Debug, Clone)]
37pub struct VisitedSet {
38    marks: Vec<u64>,
39    generation: u64,
40}
41
42impl VisitedSet {
43    /// Create a new `VisitedSet` with pre-allocated capacity for `capacity` nodes.
44    pub fn new(capacity: usize) -> Self {
45        Self {
46            marks: vec![0; capacity],
47            generation: 1,
48        }
49    }
50
51    /// Reset the visited state for all nodes in O(1) by advancing the generation.
52    #[inline]
53    pub fn clear(&mut self) {
54        self.generation = self.generation.wrapping_add(1);
55        if self.generation == 0 {
56            self.marks.fill(0);
57            self.generation = 1;
58        }
59    }
60
61    /// Grow the internal buffer if `node` would be out of range.
62    ///
63    /// Resizing uses `node + 1`; the maximum ID must permit that addition and allocation.
64    #[inline]
65    pub fn ensure_capacity(&mut self, node: usize) {
66        if node >= self.marks.len() {
67            self.marks.resize(node + 1, 0);
68        }
69    }
70
71    /// Mark `node` as visited if it has not been visited in this generation.
72    ///
73    /// Returns `true` on first visit, `false` on subsequent calls for the same node.
74    /// Resizing uses `node + 1`, with the same capacity requirements as `ensure_capacity`.
75    #[inline]
76    pub fn mark_if_new(&mut self, node: usize) -> bool {
77        if node >= self.marks.len() {
78            self.marks.resize(node + 1, 0);
79        }
80        if self.marks[node] == self.generation {
81            false
82        } else {
83            self.marks[node] = self.generation;
84            true
85        }
86    }
87
88    /// Mark a node as visited; returns `true` on its first visit in this generation.
89    #[inline]
90    pub fn visit(&mut self, node: usize) -> bool {
91        self.mark_if_new(node)
92    }
93
94    /// Mark multiple nodes as visited.
95    #[inline]
96    pub fn visit_all(&mut self, nodes: impl Iterator<Item = usize>) {
97        for node in nodes {
98            self.visit(node);
99        }
100    }
101
102    /// Return `true` if `node` has been marked in the current generation.
103    #[inline]
104    pub fn is_marked(&self, node: usize) -> bool {
105        node < self.marks.len() && self.marks[node] == self.generation
106    }
107}
108
109#[cfg(test)]
110mod tests {
111    use super::*;
112
113    #[test]
114    fn visited_set_wraparound_resets_marks() {
115        let mut vs = VisitedSet {
116            marks: vec![0; 4],
117            generation: u64::MAX,
118        };
119        vs.mark_if_new(0);
120        vs.clear();
121        assert!(vs.mark_if_new(0));
122        assert!(!vs.mark_if_new(0));
123    }
124
125    #[test]
126    fn default_is_cosine() {
127        assert_eq!(DistanceMetric::default(), DistanceMetric::Cosine);
128    }
129}