Skip to main content

mongreldb_core/index/
learned_range.rs

1//! Per-column learned (PGM) range index — serves `Condition::Range` and
2//! `Condition::RangeF64` sub-linearly for numeric columns declared
3//! `IndexKind::LearnedRange`.
4//!
5//! The run is sorted by `RowId`, not by column value, so at flush we collect
6//! `(value, row_id)`, sort by value, and build a PGM over `value → position`
7//! (parallel to the sorted `row_id` array). A range query uses the PGM to land
8//! in an ε-window, then a local binary search finds the exact `[lo,hi]` slice —
9//! `O(log segments + log ε)` instead of a full column scan.
10//!
11//! Both `i64` and `f64` columns are supported via order-preserving key encodings.
12
13use super::pgm::{LearnedIndex, PgmIndex};
14use std::collections::HashSet;
15
16/// Order-preserving encoding of `i64` into `u64` (flip the sign bit), so PGM
17/// key order matches numeric order including negatives: `MIN→0`, `-1→2⁶³-1`,
18/// `0→2⁶³`, `MAX→u64::MAX`.
19#[inline]
20fn i64_key(v: i64) -> u64 {
21    (v as u64) ^ (1u64 << 63)
22}
23
24/// Order-preserving encoding of `f64` into `u64`: positive floats map to
25/// `2⁶³..u64::MAX` (sign bit flipped), negative floats map to `0..2⁶³-1`
26/// (all bits flipped) so the total order matches IEEE-754 totalOrder.
27#[inline]
28fn f64_key(v: f64) -> u64 {
29    let bits = v.to_bits();
30    if bits & (1u64 << 63) != 0 {
31        !bits
32    } else {
33        bits ^ (1u64 << 63)
34    }
35}
36
37#[derive(Debug, Clone)]
38pub struct ColumnLearnedRange {
39    keys: Vec<u64>,    // order-preserving value keys, ascending
40    row_ids: Vec<u64>, // parallel row ids (sorted by value)
41    pgm: PgmIndex,
42}
43
44impl ColumnLearnedRange {
45    /// Build from `(value, row_id)` pairs in any order.
46    pub fn build_i64(pairs: &[(i64, u64)]) -> Self {
47        Self::build_i64_with_epsilon(pairs, 16)
48    }
49
50    pub fn build_i64_with_epsilon(pairs: &[(i64, u64)], epsilon: usize) -> Self {
51        let mut sorted: Vec<(u64, u64)> = pairs.iter().map(|(v, r)| (i64_key(*v), *r)).collect();
52        sorted.sort_unstable_by_key(|(k, _)| *k);
53        let keys: Vec<u64> = sorted.iter().map(|(k, _)| *k).collect();
54        let row_ids: Vec<u64> = sorted.iter().map(|(_, r)| *r).collect();
55        let points: Vec<(u64, usize)> = keys.iter().enumerate().map(|(i, k)| (*k, i)).collect();
56        let pgm = if points.is_empty() {
57            PgmIndex::build(&[], epsilon)
58        } else {
59            PgmIndex::build(&points, epsilon)
60        };
61        Self { keys, row_ids, pgm }
62    }
63
64    /// First position whose key `>= key`. Seeded by the PGM ε-window, then
65    /// gallops outward — the window alone only brackets duplicate runs up to
66    /// 2·ε, so longer runs need expansion (rare; keeps lookups sub-linear).
67    fn lower_bound(&self, key: u64) -> usize {
68        let n = self.keys.len();
69        if n == 0 {
70            return 0;
71        }
72        let (lo, hi) = self.pgm.predict(key);
73        let lo = lo.min(n);
74        let hi = hi.min(n).max(lo);
75        let mut idx = lo + self.keys[lo..hi].partition_point(|k| *k < key);
76        // Gallop left if the window began inside a run of keys >= search key.
77        if idx > 0 && self.keys[idx - 1] >= key {
78            let mut step = 1usize;
79            while idx >= step && self.keys[idx - step] >= key {
80                step <<= 1;
81            }
82            let start = idx - step.min(idx);
83            idx = start + self.keys[start..idx].partition_point(|k| *k < key);
84        }
85        // Gallop right if the window ended before the true lower bound.
86        if idx < n && self.keys[idx] < key {
87            let mut step = 1usize;
88            while idx + step <= n && self.keys[(idx + step).min(n) - 1] < key {
89                step <<= 1;
90            }
91            let end = (idx + step).min(n);
92            idx = idx + self.keys[idx..end].partition_point(|k| *k < key);
93        }
94        idx
95    }
96
97    /// First position whose key `> key` (galloping, same rationale).
98    fn upper_bound(&self, key: u64) -> usize {
99        let n = self.keys.len();
100        if n == 0 {
101            return 0;
102        }
103        let (lo, hi) = self.pgm.predict(key);
104        let lo = lo.min(n);
105        let hi = hi.min(n).max(lo);
106        let mut idx = lo + self.keys[lo..hi].partition_point(|k| *k <= key);
107        if idx > 0 && self.keys[idx - 1] > key {
108            let mut step = 1usize;
109            while idx >= step && self.keys[idx - step] > key {
110                step <<= 1;
111            }
112            let start = idx - step.min(idx);
113            idx = start + self.keys[start..idx].partition_point(|k| *k <= key);
114        }
115        if idx < n && self.keys[idx] <= key {
116            let mut step = 1usize;
117            while idx + step <= n && self.keys[(idx + step).min(n) - 1] <= key {
118                step <<= 1;
119            }
120            let end = (idx + step).min(n);
121            idx = idx + self.keys[idx..end].partition_point(|k| *k <= key);
122        }
123        idx
124    }
125
126    /// Row ids whose value is in `[lo, hi]` (inclusive).
127    pub fn range(&self, lo: i64, hi: i64) -> HashSet<u64> {
128        if hi < lo || self.keys.is_empty() {
129            return HashSet::new();
130        }
131        let start = self.lower_bound(i64_key(lo));
132        let end = self.upper_bound(i64_key(hi));
133        self.row_ids[start..end].iter().copied().collect()
134    }
135
136    /// Build from `(f64_value, row_id)` pairs (Phase 13.3).
137    pub fn build_f64(pairs: &[(f64, u64)]) -> Self {
138        Self::build_f64_with_epsilon(pairs, 16)
139    }
140
141    pub fn build_f64_with_epsilon(pairs: &[(f64, u64)], epsilon: usize) -> Self {
142        let mut sorted: Vec<(u64, u64)> = pairs.iter().map(|(v, r)| (f64_key(*v), *r)).collect();
143        sorted.sort_unstable_by_key(|(k, _)| *k);
144        let keys: Vec<u64> = sorted.iter().map(|(k, _)| *k).collect();
145        let row_ids: Vec<u64> = sorted.iter().map(|(_, r)| *r).collect();
146        let points: Vec<(u64, usize)> = keys.iter().enumerate().map(|(i, k)| (*k, i)).collect();
147        let pgm = if points.is_empty() {
148            PgmIndex::build(&[], epsilon)
149        } else {
150            PgmIndex::build(&points, epsilon)
151        };
152        Self { keys, row_ids, pgm }
153    }
154
155    /// Row ids whose f64 value is in `[lo, hi]` with per-bound inclusivity
156    /// (Phase 13.3).
157    pub fn range_f64(
158        &self,
159        lo: f64,
160        lo_inclusive: bool,
161        hi: f64,
162        hi_inclusive: bool,
163    ) -> HashSet<u64> {
164        if self.keys.is_empty() {
165            return HashSet::new();
166        }
167        // Convert to key space. For inclusive bounds, use the value's key
168        // directly (it maps to the exact position). For exclusive bounds,
169        // nudge inward by ±1 in key space (which is the next representable
170        // value in the order-preserving encoding).
171        let lo_key = f64_key(lo);
172        let hi_key = f64_key(hi);
173        let (start, end) = if hi < lo {
174            return HashSet::new();
175        } else if lo_inclusive && hi_inclusive {
176            (self.lower_bound(lo_key), self.upper_bound(hi_key))
177        } else if lo_inclusive {
178            // hi exclusive
179            (self.lower_bound(lo_key), self.lower_bound(hi_key))
180        } else if hi_inclusive {
181            // lo exclusive
182            (self.upper_bound(lo_key), self.upper_bound(hi_key))
183        } else {
184            (self.upper_bound(lo_key), self.lower_bound(hi_key))
185        };
186        self.row_ids[start..end].iter().copied().collect()
187    }
188
189    /// Snapshot `(value_key, row_id, pgm segments/epsilon)` for checkpointing.
190    pub fn snapshot(&self) -> ColumnLearnedRangeSnapshot {
191        ColumnLearnedRangeSnapshot {
192            keys: self.keys.clone(),
193            row_ids: self.row_ids.clone(),
194            pgm: self.pgm.clone(),
195        }
196    }
197
198    /// Rebuild from a snapshot produced by [`ColumnLearnedRange::snapshot`].
199    pub fn from_snapshot(snap: ColumnLearnedRangeSnapshot) -> Self {
200        Self {
201            keys: snap.keys,
202            row_ids: snap.row_ids,
203            pgm: snap.pgm,
204        }
205    }
206}
207
208/// Serializable snapshot of a [`ColumnLearnedRange`] (PGM is already serde).
209#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
210pub struct ColumnLearnedRangeSnapshot {
211    pub keys: Vec<u64>,
212    pub row_ids: Vec<u64>,
213    pub pgm: PgmIndex,
214}
215
216#[cfg(test)]
217mod tests {
218    use super::*;
219
220    #[test]
221    fn range_returns_exact_slice() {
222        // values out of row_id order; duplicates present.
223        let pairs = vec![
224            (100i64, 0u64),
225            (-5, 1),
226            (50, 2),
227            (100, 3),
228            (1000, 4),
229            (-5, 5),
230        ];
231        let idx = ColumnLearnedRange::build_i64(&pairs);
232        // sorted by value: -5(r1,r5), 50(r2), 100(r0,r3), 1000(r4)
233        let r = idx.range(50, 100);
234        assert_eq!(r, [2, 0, 3].into_iter().collect::<HashSet<_>>());
235        let all = idx.range(i64::MIN, i64::MAX);
236        assert_eq!(all.len(), 6);
237        let none = idx.range(200, 300);
238        assert!(none.is_empty());
239        let negs = idx.range(i64::MIN, -5);
240        assert_eq!(negs, [1, 5].into_iter().collect::<HashSet<_>>());
241    }
242}