Skip to main content

formualizer_eval/engine/
interval_tree.rs

1use std::collections::{BTreeMap, HashSet};
2use std::ops::ControlFlow;
3
4/// Custom interval tree optimized for spreadsheet cell indexing.
5///
6/// ## Design decisions:
7///
8/// 1. **Point intervals are the common case** - Most cells are single points [r,r] or [c,c]
9/// 2. **Sparse data** - Even million-row sheets typically have <10K cells
10/// 3. **Batch updates** - During shifts, we update many intervals at once
11/// 4. **Small value sets** - Each interval maps to a small set of VertexIds
12///
13/// ## Implementation:
14///
15/// Uses an augmented BST where each node stores:
16/// - Interval [low, high]
17/// - Max endpoint in subtree (for efficient pruning)
18/// - Value set (HashSet<VertexId>)
19///
20/// This is simpler than generic interval trees because we optimize for our specific use case.
21
22#[derive(Debug, Clone)]
23struct IntervalNode<T: Clone + Eq + std::hash::Hash> {
24    high: u32,
25    values: HashSet<T>,
26}
27
28/// B-Tree based implementation of the interval index.
29#[derive(Debug, Clone)]
30pub struct IntervalTree<T: Clone + Eq + std::hash::Hash> {
31    /// Maps low coordinate to a set of intervals/values starting there.
32    /// Internal storage uses IntervalNode, NOT Entry.
33    map: BTreeMap<u32, Vec<IntervalNode<T>>>,
34    size: usize,
35}
36
37impl<T: Clone + Eq + std::hash::Hash> Default for IntervalTree<T> {
38    fn default() -> Self {
39        Self::new()
40    }
41}
42
43impl<T: Clone + Eq + std::hash::Hash> IntervalTree<T> {
44    pub fn new() -> Self {
45        Self {
46            map: BTreeMap::new(),
47            size: 0,
48        }
49    }
50
51    pub fn len(&self) -> usize {
52        self.size
53    }
54
55    pub fn is_empty(&self) -> bool {
56        self.size == 0
57    }
58
59    /// Get a mutable reference to the values for an exact interval match.
60    /// Required by the Entry API.
61    pub fn get_mut(&mut self, low: u32, high: u32) -> Option<&mut HashSet<T>> {
62        self.map.get_mut(&low).and_then(|nodes| {
63            nodes
64                .iter_mut()
65                .find(|n| n.high == high)
66                .map(|n| &mut n.values)
67        })
68    }
69
70    /// Insert a value for the given interval [low, high]
71    pub fn insert(&mut self, low: u32, high: u32, value: T) {
72        let entries = self.map.entry(low).or_default();
73
74        if let Some(node) = entries.iter_mut().find(|n| n.high == high) {
75            node.values.insert(value);
76        } else {
77            let mut values = HashSet::new();
78            values.insert(value);
79            entries.push(IntervalNode { high, values });
80            self.size += 1;
81        }
82    }
83
84    pub fn query(&self, q_low: u32, q_high: u32) -> Vec<(u32, u32, HashSet<T>)> {
85        let mut results = Vec::new();
86        for (&low, nodes) in self.map.range(..=q_high) {
87            for node in nodes {
88                if node.high >= q_low {
89                    results.push((low, node.high, node.values.clone()));
90                }
91            }
92        }
93        results
94    }
95
96    /// Returns the values stored at the exact point interval `[point, point]`.
97    ///
98    /// Unlike [`Self::query`], this performs a direct B-tree lookup and therefore
99    /// does not inspect intervals whose lower bound precedes `point`.
100    pub(crate) fn point_values(&self, point: u32) -> Option<&HashSet<T>> {
101        self.map.get(&point).and_then(|nodes| {
102            nodes
103                .iter()
104                .find(|node| node.high == point)
105                .map(|node| &node.values)
106        })
107    }
108
109    /// Visits point intervals whose coordinate is in the inclusive query range.
110    ///
111    /// This is a specialized path for indexes that contain only `[x, x]`
112    /// intervals. General overlapping-range callers must continue to use
113    /// [`Self::query`] or [`Self::visit_query`]. `None` is emitted once for each
114    /// point-interval node inspected so callers can account deterministic visits.
115    pub(crate) fn visit_point_intervals(
116        &self,
117        q_low: u32,
118        q_high: u32,
119        mut visitor: impl FnMut(Option<&T>) -> ControlFlow<()>,
120    ) -> ControlFlow<()> {
121        if q_low > q_high {
122            return ControlFlow::Continue(());
123        }
124        for (&low, nodes) in self.map.range(q_low..=q_high) {
125            for node in nodes {
126                if node.high != low {
127                    continue;
128                }
129                visitor(None)?;
130                for value in &node.values {
131                    visitor(Some(value))?;
132                }
133            }
134        }
135        ControlFlow::Continue(())
136    }
137
138    /// Counts point-interval nodes and values in an inclusive coordinate range.
139    /// This has the same point-only caller contract as [`Self::visit_point_intervals`].
140    pub(crate) fn point_interval_stats(&self, q_low: u32, q_high: u32) -> (usize, usize) {
141        if q_low > q_high {
142            return (0, 0);
143        }
144        let mut node_count = 0usize;
145        let mut value_count = 0usize;
146        for (&low, nodes) in self.map.range(q_low..=q_high) {
147            for node in nodes {
148                if node.high == low {
149                    node_count = node_count.saturating_add(1);
150                    value_count = value_count.saturating_add(node.values.len());
151                }
152            }
153        }
154        (node_count, value_count)
155    }
156
157    /// Visits matching values without cloning or materializing the query.
158    /// Returning `Break` from the visitor stops traversal immediately.
159    pub(crate) fn visit_query(
160        &self,
161        q_low: u32,
162        q_high: u32,
163        mut visitor: impl FnMut(Option<&T>) -> ControlFlow<()>,
164    ) -> ControlFlow<()> {
165        for (_low, nodes) in self.map.range(..=q_high) {
166            for node in nodes {
167                visitor(None)?;
168                if node.high < q_low {
169                    continue;
170                }
171                for value in &node.values {
172                    visitor(Some(value))?;
173                }
174            }
175        }
176        ControlFlow::Continue(())
177    }
178
179    pub(crate) fn estimated_heap_bytes(&self) -> Option<usize> {
180        const NODE_ALLOCATION_OVERHEAD: usize = 3 * std::mem::size_of::<usize>();
181        let mut bytes = 0usize;
182        for nodes in self.map.values() {
183            bytes = bytes.checked_add(
184                std::mem::size_of::<u32>()
185                    .checked_add(std::mem::size_of::<Vec<IntervalNode<T>>>())?
186                    .checked_add(NODE_ALLOCATION_OVERHEAD)?,
187            )?;
188            bytes = bytes.checked_add(
189                nodes
190                    .capacity()
191                    .checked_mul(std::mem::size_of::<IntervalNode<T>>())?,
192            )?;
193            for node in nodes {
194                bytes = bytes.checked_add(node.values.capacity().checked_mul(
195                    std::mem::size_of::<T>().checked_add(std::mem::size_of::<usize>())?,
196                )?)?;
197            }
198        }
199        Some(bytes)
200    }
201
202    pub fn remove(&mut self, low: u32, high: u32, value: &T) -> bool {
203        if let Some(nodes) = self.map.get_mut(&low)
204            && let Some(node) = nodes.iter_mut().find(|n| n.high == high)
205        {
206            let removed = node.values.remove(value);
207
208            if removed && node.values.is_empty() {
209                nodes.retain(|n| n.high != high);
210                self.size -= 1;
211                if nodes.is_empty() {
212                    self.map.remove(&low);
213                }
214            }
215            return removed;
216        }
217        false
218    }
219
220    /// Keep only the values satisfying `keep` (empty intervals are
221    /// dropped), in one pass.
222    pub(crate) fn retain_values(&mut self, mut keep: impl FnMut(&T) -> bool) {
223        let mut size = 0;
224        self.map.retain(|_, nodes| {
225            nodes.retain_mut(|n| {
226                n.values.retain(&mut keep);
227                if n.values.is_empty() {
228                    false
229                } else {
230                    // Give capacity back only when most of it is unused.
231                    if n.values.len() * 4 < n.values.capacity() {
232                        n.values.shrink_to_fit();
233                    }
234                    true
235                }
236            });
237            size += nodes.len();
238            !nodes.is_empty()
239        });
240        self.size = size;
241    }
242
243    pub fn entry(&mut self, low: u32, high: u32) -> BTreeEntry<'_, T> {
244        BTreeEntry {
245            tree: self,
246            low,
247            high,
248        }
249    }
250
251    /// Bulk build optimization for a collection of point intervals [x,x].
252    pub fn bulk_build_points(&mut self, mut items: Vec<(u32, HashSet<T>)>) {
253        if !self.is_empty() {
254            // Fallback: incremental insert to preserve existing nodes
255            for (coord, set) in items {
256                for val in set {
257                    self.insert(coord, coord, val);
258                }
259            }
260            return;
261        }
262
263        if items.is_empty() {
264            return;
265        }
266
267        // 1. Sort by coordinate
268        items.sort_by_key(|(k, _)| *k);
269
270        // 2. Process items. BTreeMap handles the balancing (O(log N)).
271        for (coord, set) in items {
272            let entries = self.map.entry(coord).or_default();
273
274            // Since this is specifically for point intervals, check if [coord, coord] exists
275            if let Some(node) = entries.iter_mut().find(|n| n.high == coord) {
276                node.values.extend(set);
277            } else {
278                entries.push(IntervalNode {
279                    high: coord,
280                    values: set,
281                });
282                self.size += 1;
283            }
284        }
285    }
286}
287
288pub struct BTreeEntry<'a, T: Clone + Eq + std::hash::Hash> {
289    tree: &'a mut IntervalTree<T>,
290    low: u32,
291    high: u32,
292}
293
294impl<'a, T: Clone + Eq + std::hash::Hash> BTreeEntry<'a, T> {
295    pub fn or_insert_with<F>(self, f: F) -> &'a mut HashSet<T>
296    where
297        F: FnOnce() -> HashSet<T>,
298    {
299        if self.tree.get_mut(self.low, self.high).is_none() {
300            let values = f();
301            let entries = self.tree.map.entry(self.low).or_default();
302            entries.push(IntervalNode {
303                high: self.high,
304                values,
305            });
306            self.tree.size += 1;
307        }
308        self.tree.get_mut(self.low, self.high).unwrap()
309    }
310}
311
312#[cfg(test)]
313mod tests {
314    use super::*;
315
316    #[test]
317    fn test_insert_and_query_point_interval() {
318        let mut tree = IntervalTree::new();
319        tree.insert(5, 5, 100);
320
321        let results = tree.query(5, 5);
322        assert_eq!(results.len(), 1);
323        assert_eq!(results[0].0, 5);
324        assert_eq!(results[0].1, 5);
325        assert!(results[0].2.contains(&100));
326    }
327
328    #[test]
329    fn test_insert_and_query_range() {
330        let mut tree = IntervalTree::new();
331        tree.insert(10, 20, 1);
332        tree.insert(15, 25, 2);
333        tree.insert(30, 40, 3);
334
335        // Query overlapping with first two intervals
336        let results = tree.query(12, 22);
337        assert_eq!(results.len(), 2);
338
339        // Query overlapping with only the third interval
340        let results = tree.query(35, 45);
341        assert_eq!(results.len(), 1);
342        assert!(results[0].2.contains(&3));
343    }
344
345    #[test]
346    fn point_interval_visit_does_not_scan_coordinate_prefixes() {
347        let mut tree = IntervalTree::new();
348        for coordinate in 0..10_000 {
349            tree.insert(coordinate, coordinate, coordinate);
350        }
351
352        let mut nodes = 0;
353        let mut values = Vec::new();
354        let result = tree.visit_point_intervals(9_999, 9_999, |entry| {
355            match entry {
356                None => nodes += 1,
357                Some(value) => values.push(*value),
358            }
359            ControlFlow::Continue(())
360        });
361
362        assert_eq!(result, ControlFlow::Continue(()));
363        assert_eq!(nodes, 1);
364        assert_eq!(values, vec![9_999]);
365    }
366
367    #[test]
368    fn point_interval_path_does_not_change_general_overlap_queries() {
369        let mut tree = IntervalTree::new();
370        tree.insert(1, 100, "range");
371        tree.insert(75, 75, "point");
372
373        let general = tree.query(75, 75);
374        assert_eq!(general.len(), 2);
375        assert!(general.iter().any(|entry| entry.2.contains("range")));
376        assert!(general.iter().any(|entry| entry.2.contains("point")));
377
378        let mut point_values = Vec::new();
379        let _ = tree.visit_point_intervals(75, 75, |entry| {
380            if let Some(value) = entry {
381                point_values.push(*value);
382            }
383            ControlFlow::Continue(())
384        });
385        assert_eq!(point_values, vec!["point"]);
386    }
387
388    #[test]
389    fn test_remove_value() {
390        let mut tree = IntervalTree::new();
391        tree.insert(5, 5, 100);
392        tree.insert(5, 5, 200);
393
394        assert_eq!(tree.query(5, 5).len(), 1);
395        assert_eq!(tree.query(5, 5)[0].2.len(), 2);
396
397        tree.remove(5, 5, &100);
398
399        let results = tree.query(5, 5);
400        assert_eq!(results.len(), 1);
401        assert_eq!(results[0].2.len(), 1);
402        assert!(results[0].2.contains(&200));
403    }
404
405    #[test]
406    fn test_entry_api() {
407        let mut tree: IntervalTree<i32> = IntervalTree::new();
408
409        tree.entry(10, 10).or_insert_with(HashSet::new).insert(42);
410
411        tree.entry(10, 10).or_insert_with(HashSet::new).insert(43);
412
413        let results = tree.query(10, 10);
414        assert_eq!(results.len(), 1);
415        assert_eq!(results[0].2.len(), 2);
416        assert!(results[0].2.contains(&42));
417        assert!(results[0].2.contains(&43));
418    }
419
420    #[test]
421    fn test_large_sparse_tree() {
422        let mut tree = IntervalTree::new();
423
424        // Simulate sparse spreadsheet
425        for i in (0..1_000_000).step_by(10000) {
426            tree.insert(i, i, i as i32);
427        }
428
429        assert_eq!(tree.len(), 100);
430
431        // Query for high rows
432        let results = tree.query(500_000, u32::MAX);
433        assert_eq!(results.len(), 50);
434    }
435
436    #[test]
437    fn test_entry_recursion_bug() {
438        let mut tree: IntervalTree<u32> = IntervalTree::new();
439
440        // The bug happens when we insert a value, then use entry()
441        // on a coordinate that would be a child of that value.
442        let count: u32 = 5000;
443        for i in 0..count {
444            tree.entry(i, i).or_insert_with(HashSet::new);
445        }
446
447        assert_eq!(tree.len(), count as usize);
448    }
449
450    #[test]
451    fn test_complex_overlaps() {
452        let mut tree = IntervalTree::new();
453        // Nested intervals
454        tree.insert(10, 100, "A");
455        tree.insert(20, 50, "B");
456        tree.insert(30, 40, "C");
457
458        // Partially overlapping
459        tree.insert(5, 15, "D");
460        tree.insert(95, 105, "E");
461
462        // Query for the very middle
463        let results = tree.query(35, 35);
464        assert_eq!(results.len(), 3); // Should hit A, B, and C
465
466        // Query for a range that only hits the "tail" of the large interval and the "head" of the end interval
467        let results = tree.query(98, 102);
468        assert_eq!(results.len(), 2); // Should hit A and E
469    }
470
471    #[test]
472    fn test_multiple_values_and_size() {
473        let mut tree = IntervalTree::new();
474
475        // Insert same interval twice with different values
476        tree.insert(10, 10, "val1");
477        tree.insert(10, 10, "val2");
478        assert_eq!(tree.len(), 1); // Size should only count unique intervals
479
480        // Insert same value twice
481        tree.insert(10, 10, "val1");
482        assert_eq!(tree.len(), 1);
483        let results = tree.query(10, 10);
484        assert_eq!(results[0].2.len(), 2); // HashSet handles the duplicate value "val1"
485    }
486
487    #[test]
488    fn test_remove_edge_cases() {
489        let mut tree = IntervalTree::new();
490        tree.insert(10, 20, "A");
491
492        // Try to remove a value that isn't there
493        let removed = tree.remove(10, 20, &"B");
494        assert!(!removed);
495        assert_eq!(tree.query(10, 20)[0].2.len(), 1);
496
497        // Try to remove from an interval that doesn't exist
498        let removed = tree.remove(99, 100, &"A");
499        assert!(!removed);
500    }
501
502    #[test]
503    fn test_bulk_build_consistency() {
504        let mut incremental_tree = IntervalTree::new();
505        let mut bulk_tree = IntervalTree::new();
506
507        let data: Vec<(u32, HashSet<&str>)> = vec![
508            (10, vec!["A", "B"].into_iter().collect()),
509            (20, vec!["C"].into_iter().collect()),
510            (5, vec!["D"].into_iter().collect()),
511        ];
512
513        // Build incrementally
514        for (coord, values) in &data {
515            for val in values {
516                incremental_tree.insert(*coord, *coord, *val);
517            }
518        }
519
520        // Build using bulk
521        bulk_tree.bulk_build_points(data.clone());
522
523        // Compare results
524        assert_eq!(incremental_tree.len(), bulk_tree.len());
525        assert_eq!(incremental_tree.query(0, 100), bulk_tree.query(0, 100));
526    }
527
528    #[test]
529    fn test_query_stack_safety() {
530        let mut tree = IntervalTree::new();
531        let count = 10_000;
532
533        // Create a deep right-leaning tree
534        for i in 0..count {
535            tree.insert(i, i, i);
536        }
537
538        // Query the very end of the tree
539        // If this causes a SIGABRT, it means query_node() must be made iterative
540        let results = tree.query(count - 1, count - 1);
541        assert_eq!(results.len(), 1);
542    }
543
544    #[test]
545    fn test_empty_and_boundaries() {
546        let mut tree: IntervalTree<i32> = IntervalTree::new();
547
548        assert!(tree.is_empty());
549        assert_eq!(tree.query(0, 100).len(), 0);
550        assert!(!tree.remove(0, 0, &1));
551
552        // Test a query that "misses" everything
553        tree.insert(50, 60, 1);
554        assert_eq!(tree.query(0, 49).len(), 0);
555        assert_eq!(tree.query(61, 100).len(), 0);
556    }
557
558    #[test]
559    fn test_multi_value_interval_size_tracking() {
560        let mut tree = IntervalTree::new();
561        let iv = (10, 20);
562
563        // 1. Insert two values for the same interval
564        // Destructure the tuple into low (iv.0) and high (iv.1)
565        tree.insert(iv.0, iv.1, "A");
566        tree.insert(iv.0, iv.1, "B");
567        assert_eq!(tree.len(), 1, "Should be 1 unique interval");
568
569        // 2. Remove first value - pass as reference &"A"
570        assert!(tree.remove(iv.0, iv.1, &"A"));
571        assert_eq!(
572            tree.len(),
573            1,
574            "Should still be 1 interval after partial removal"
575        );
576
577        // 3. Remove second value - size should now be 0
578        assert!(tree.remove(iv.0, iv.1, &"B"));
579        assert_eq!(tree.len(), 0, "Should be 0 after last value removed");
580    }
581}