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    pub fn entry(&mut self, low: u32, high: u32) -> BTreeEntry<'_, T> {
221        BTreeEntry {
222            tree: self,
223            low,
224            high,
225        }
226    }
227
228    /// Bulk build optimization for a collection of point intervals [x,x].
229    pub fn bulk_build_points(&mut self, mut items: Vec<(u32, HashSet<T>)>) {
230        if !self.is_empty() {
231            // Fallback: incremental insert to preserve existing nodes
232            for (coord, set) in items {
233                for val in set {
234                    self.insert(coord, coord, val);
235                }
236            }
237            return;
238        }
239
240        if items.is_empty() {
241            return;
242        }
243
244        // 1. Sort by coordinate
245        items.sort_by_key(|(k, _)| *k);
246
247        // 2. Process items. BTreeMap handles the balancing (O(log N)).
248        for (coord, set) in items {
249            let entries = self.map.entry(coord).or_default();
250
251            // Since this is specifically for point intervals, check if [coord, coord] exists
252            if let Some(node) = entries.iter_mut().find(|n| n.high == coord) {
253                node.values.extend(set);
254            } else {
255                entries.push(IntervalNode {
256                    high: coord,
257                    values: set,
258                });
259                self.size += 1;
260            }
261        }
262    }
263}
264
265pub struct BTreeEntry<'a, T: Clone + Eq + std::hash::Hash> {
266    tree: &'a mut IntervalTree<T>,
267    low: u32,
268    high: u32,
269}
270
271impl<'a, T: Clone + Eq + std::hash::Hash> BTreeEntry<'a, T> {
272    pub fn or_insert_with<F>(self, f: F) -> &'a mut HashSet<T>
273    where
274        F: FnOnce() -> HashSet<T>,
275    {
276        if self.tree.get_mut(self.low, self.high).is_none() {
277            let values = f();
278            let entries = self.tree.map.entry(self.low).or_default();
279            entries.push(IntervalNode {
280                high: self.high,
281                values,
282            });
283            self.tree.size += 1;
284        }
285        self.tree.get_mut(self.low, self.high).unwrap()
286    }
287}
288
289#[cfg(test)]
290mod tests {
291    use super::*;
292
293    #[test]
294    fn test_insert_and_query_point_interval() {
295        let mut tree = IntervalTree::new();
296        tree.insert(5, 5, 100);
297
298        let results = tree.query(5, 5);
299        assert_eq!(results.len(), 1);
300        assert_eq!(results[0].0, 5);
301        assert_eq!(results[0].1, 5);
302        assert!(results[0].2.contains(&100));
303    }
304
305    #[test]
306    fn test_insert_and_query_range() {
307        let mut tree = IntervalTree::new();
308        tree.insert(10, 20, 1);
309        tree.insert(15, 25, 2);
310        tree.insert(30, 40, 3);
311
312        // Query overlapping with first two intervals
313        let results = tree.query(12, 22);
314        assert_eq!(results.len(), 2);
315
316        // Query overlapping with only the third interval
317        let results = tree.query(35, 45);
318        assert_eq!(results.len(), 1);
319        assert!(results[0].2.contains(&3));
320    }
321
322    #[test]
323    fn point_interval_visit_does_not_scan_coordinate_prefixes() {
324        let mut tree = IntervalTree::new();
325        for coordinate in 0..10_000 {
326            tree.insert(coordinate, coordinate, coordinate);
327        }
328
329        let mut nodes = 0;
330        let mut values = Vec::new();
331        let result = tree.visit_point_intervals(9_999, 9_999, |entry| {
332            match entry {
333                None => nodes += 1,
334                Some(value) => values.push(*value),
335            }
336            ControlFlow::Continue(())
337        });
338
339        assert_eq!(result, ControlFlow::Continue(()));
340        assert_eq!(nodes, 1);
341        assert_eq!(values, vec![9_999]);
342    }
343
344    #[test]
345    fn point_interval_path_does_not_change_general_overlap_queries() {
346        let mut tree = IntervalTree::new();
347        tree.insert(1, 100, "range");
348        tree.insert(75, 75, "point");
349
350        let general = tree.query(75, 75);
351        assert_eq!(general.len(), 2);
352        assert!(general.iter().any(|entry| entry.2.contains("range")));
353        assert!(general.iter().any(|entry| entry.2.contains("point")));
354
355        let mut point_values = Vec::new();
356        let _ = tree.visit_point_intervals(75, 75, |entry| {
357            if let Some(value) = entry {
358                point_values.push(*value);
359            }
360            ControlFlow::Continue(())
361        });
362        assert_eq!(point_values, vec!["point"]);
363    }
364
365    #[test]
366    fn test_remove_value() {
367        let mut tree = IntervalTree::new();
368        tree.insert(5, 5, 100);
369        tree.insert(5, 5, 200);
370
371        assert_eq!(tree.query(5, 5).len(), 1);
372        assert_eq!(tree.query(5, 5)[0].2.len(), 2);
373
374        tree.remove(5, 5, &100);
375
376        let results = tree.query(5, 5);
377        assert_eq!(results.len(), 1);
378        assert_eq!(results[0].2.len(), 1);
379        assert!(results[0].2.contains(&200));
380    }
381
382    #[test]
383    fn test_entry_api() {
384        let mut tree: IntervalTree<i32> = IntervalTree::new();
385
386        tree.entry(10, 10).or_insert_with(HashSet::new).insert(42);
387
388        tree.entry(10, 10).or_insert_with(HashSet::new).insert(43);
389
390        let results = tree.query(10, 10);
391        assert_eq!(results.len(), 1);
392        assert_eq!(results[0].2.len(), 2);
393        assert!(results[0].2.contains(&42));
394        assert!(results[0].2.contains(&43));
395    }
396
397    #[test]
398    fn test_large_sparse_tree() {
399        let mut tree = IntervalTree::new();
400
401        // Simulate sparse spreadsheet
402        for i in (0..1_000_000).step_by(10000) {
403            tree.insert(i, i, i as i32);
404        }
405
406        assert_eq!(tree.len(), 100);
407
408        // Query for high rows
409        let results = tree.query(500_000, u32::MAX);
410        assert_eq!(results.len(), 50);
411    }
412
413    #[test]
414    fn test_entry_recursion_bug() {
415        let mut tree: IntervalTree<u32> = IntervalTree::new();
416
417        // The bug happens when we insert a value, then use entry()
418        // on a coordinate that would be a child of that value.
419        let count: u32 = 5000;
420        for i in 0..count {
421            tree.entry(i, i).or_insert_with(HashSet::new);
422        }
423
424        assert_eq!(tree.len(), count as usize);
425    }
426
427    #[test]
428    fn test_complex_overlaps() {
429        let mut tree = IntervalTree::new();
430        // Nested intervals
431        tree.insert(10, 100, "A");
432        tree.insert(20, 50, "B");
433        tree.insert(30, 40, "C");
434
435        // Partially overlapping
436        tree.insert(5, 15, "D");
437        tree.insert(95, 105, "E");
438
439        // Query for the very middle
440        let results = tree.query(35, 35);
441        assert_eq!(results.len(), 3); // Should hit A, B, and C
442
443        // Query for a range that only hits the "tail" of the large interval and the "head" of the end interval
444        let results = tree.query(98, 102);
445        assert_eq!(results.len(), 2); // Should hit A and E
446    }
447
448    #[test]
449    fn test_multiple_values_and_size() {
450        let mut tree = IntervalTree::new();
451
452        // Insert same interval twice with different values
453        tree.insert(10, 10, "val1");
454        tree.insert(10, 10, "val2");
455        assert_eq!(tree.len(), 1); // Size should only count unique intervals
456
457        // Insert same value twice
458        tree.insert(10, 10, "val1");
459        assert_eq!(tree.len(), 1);
460        let results = tree.query(10, 10);
461        assert_eq!(results[0].2.len(), 2); // HashSet handles the duplicate value "val1"
462    }
463
464    #[test]
465    fn test_remove_edge_cases() {
466        let mut tree = IntervalTree::new();
467        tree.insert(10, 20, "A");
468
469        // Try to remove a value that isn't there
470        let removed = tree.remove(10, 20, &"B");
471        assert!(!removed);
472        assert_eq!(tree.query(10, 20)[0].2.len(), 1);
473
474        // Try to remove from an interval that doesn't exist
475        let removed = tree.remove(99, 100, &"A");
476        assert!(!removed);
477    }
478
479    #[test]
480    fn test_bulk_build_consistency() {
481        let mut incremental_tree = IntervalTree::new();
482        let mut bulk_tree = IntervalTree::new();
483
484        let data: Vec<(u32, HashSet<&str>)> = vec![
485            (10, vec!["A", "B"].into_iter().collect()),
486            (20, vec!["C"].into_iter().collect()),
487            (5, vec!["D"].into_iter().collect()),
488        ];
489
490        // Build incrementally
491        for (coord, values) in &data {
492            for val in values {
493                incremental_tree.insert(*coord, *coord, *val);
494            }
495        }
496
497        // Build using bulk
498        bulk_tree.bulk_build_points(data.clone());
499
500        // Compare results
501        assert_eq!(incremental_tree.len(), bulk_tree.len());
502        assert_eq!(incremental_tree.query(0, 100), bulk_tree.query(0, 100));
503    }
504
505    #[test]
506    fn test_query_stack_safety() {
507        let mut tree = IntervalTree::new();
508        let count = 10_000;
509
510        // Create a deep right-leaning tree
511        for i in 0..count {
512            tree.insert(i, i, i);
513        }
514
515        // Query the very end of the tree
516        // If this causes a SIGABRT, it means query_node() must be made iterative
517        let results = tree.query(count - 1, count - 1);
518        assert_eq!(results.len(), 1);
519    }
520
521    #[test]
522    fn test_empty_and_boundaries() {
523        let mut tree: IntervalTree<i32> = IntervalTree::new();
524
525        assert!(tree.is_empty());
526        assert_eq!(tree.query(0, 100).len(), 0);
527        assert!(!tree.remove(0, 0, &1));
528
529        // Test a query that "misses" everything
530        tree.insert(50, 60, 1);
531        assert_eq!(tree.query(0, 49).len(), 0);
532        assert_eq!(tree.query(61, 100).len(), 0);
533    }
534
535    #[test]
536    fn test_multi_value_interval_size_tracking() {
537        let mut tree = IntervalTree::new();
538        let iv = (10, 20);
539
540        // 1. Insert two values for the same interval
541        // Destructure the tuple into low (iv.0) and high (iv.1)
542        tree.insert(iv.0, iv.1, "A");
543        tree.insert(iv.0, iv.1, "B");
544        assert_eq!(tree.len(), 1, "Should be 1 unique interval");
545
546        // 2. Remove first value - pass as reference &"A"
547        assert!(tree.remove(iv.0, iv.1, &"A"));
548        assert_eq!(
549            tree.len(),
550            1,
551            "Should still be 1 interval after partial removal"
552        );
553
554        // 3. Remove second value - size should now be 0
555        assert!(tree.remove(iv.0, iv.1, &"B"));
556        assert_eq!(tree.len(), 0, "Should be 0 after last value removed");
557    }
558}