1use super::pgm::{LearnedIndex, PgmIndex};
14use std::collections::HashSet;
15
16#[inline]
20fn i64_key(v: i64) -> u64 {
21 (v as u64) ^ (1u64 << 63)
22}
23
24#[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>, row_ids: Vec<u64>, pgm: PgmIndex,
42}
43
44impl ColumnLearnedRange {
45 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 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 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 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 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 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 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 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 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 (self.lower_bound(lo_key), self.lower_bound(hi_key))
180 } else if hi_inclusive {
181 (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 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 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#[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 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 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}