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    /// Like [`Self::point_interval_stats`] but stops once more than `cap`
158    /// entries (point nodes plus values) have been seen: `None` means the
159    /// range holds more than `cap`, at a cost of O(cap) instead of
160    /// O(range).
161    pub(crate) fn point_interval_size_capped(
162        &self,
163        q_low: u32,
164        q_high: u32,
165        cap: usize,
166    ) -> Option<usize> {
167        if q_low > q_high {
168            return Some(0);
169        }
170        let mut seen = 0usize;
171        for (&low, nodes) in self.map.range(q_low..=q_high) {
172            for node in nodes {
173                if node.high == low {
174                    seen = seen.saturating_add(1 + node.values.len());
175                    if seen > cap {
176                        return None;
177                    }
178                }
179            }
180        }
181        Some(seen)
182    }
183
184    /// Visits matching values without cloning or materializing the query.
185    /// Returning `Break` from the visitor stops traversal immediately.
186    pub(crate) fn visit_query(
187        &self,
188        q_low: u32,
189        q_high: u32,
190        mut visitor: impl FnMut(Option<&T>) -> ControlFlow<()>,
191    ) -> ControlFlow<()> {
192        for (_low, nodes) in self.map.range(..=q_high) {
193            for node in nodes {
194                visitor(None)?;
195                if node.high < q_low {
196                    continue;
197                }
198                for value in &node.values {
199                    visitor(Some(value))?;
200                }
201            }
202        }
203        ControlFlow::Continue(())
204    }
205
206    pub(crate) fn estimated_heap_bytes(&self) -> Option<usize> {
207        const NODE_ALLOCATION_OVERHEAD: usize = 3 * std::mem::size_of::<usize>();
208        let mut bytes = 0usize;
209        for nodes in self.map.values() {
210            bytes = bytes.checked_add(
211                std::mem::size_of::<u32>()
212                    .checked_add(std::mem::size_of::<Vec<IntervalNode<T>>>())?
213                    .checked_add(NODE_ALLOCATION_OVERHEAD)?,
214            )?;
215            bytes = bytes.checked_add(
216                nodes
217                    .capacity()
218                    .checked_mul(std::mem::size_of::<IntervalNode<T>>())?,
219            )?;
220            for node in nodes {
221                bytes = bytes.checked_add(node.values.capacity().checked_mul(
222                    std::mem::size_of::<T>().checked_add(std::mem::size_of::<usize>())?,
223                )?)?;
224            }
225        }
226        Some(bytes)
227    }
228
229    pub fn remove(&mut self, low: u32, high: u32, value: &T) -> bool {
230        if let Some(nodes) = self.map.get_mut(&low)
231            && let Some(node) = nodes.iter_mut().find(|n| n.high == high)
232        {
233            let removed = node.values.remove(value);
234
235            if removed && node.values.is_empty() {
236                nodes.retain(|n| n.high != high);
237                self.size -= 1;
238                if nodes.is_empty() {
239                    self.map.remove(&low);
240                }
241            }
242            return removed;
243        }
244        false
245    }
246
247    /// Keep only the values satisfying `keep` (empty intervals are
248    /// dropped), in one pass.
249    pub(crate) fn retain_values(&mut self, mut keep: impl FnMut(&T) -> bool) {
250        let mut size = 0;
251        self.map.retain(|_, nodes| {
252            nodes.retain_mut(|n| {
253                n.values.retain(&mut keep);
254                if n.values.is_empty() {
255                    false
256                } else {
257                    // Give capacity back only when most of it is unused.
258                    if n.values.len() * 4 < n.values.capacity() {
259                        n.values.shrink_to_fit();
260                    }
261                    true
262                }
263            });
264            size += nodes.len();
265            !nodes.is_empty()
266        });
267        self.size = size;
268    }
269
270    pub fn entry(&mut self, low: u32, high: u32) -> BTreeEntry<'_, T> {
271        BTreeEntry {
272            tree: self,
273            low,
274            high,
275        }
276    }
277
278    /// Bulk build optimization for a collection of point intervals [x,x].
279    pub fn bulk_build_points(&mut self, mut items: Vec<(u32, HashSet<T>)>) {
280        if !self.is_empty() {
281            // Fallback: incremental insert to preserve existing nodes
282            for (coord, set) in items {
283                for val in set {
284                    self.insert(coord, coord, val);
285                }
286            }
287            return;
288        }
289
290        if items.is_empty() {
291            return;
292        }
293
294        // 1. Sort by coordinate
295        items.sort_by_key(|(k, _)| *k);
296
297        // 2. Process items. BTreeMap handles the balancing (O(log N)).
298        for (coord, set) in items {
299            let entries = self.map.entry(coord).or_default();
300
301            // Since this is specifically for point intervals, check if [coord, coord] exists
302            if let Some(node) = entries.iter_mut().find(|n| n.high == coord) {
303                node.values.extend(set);
304            } else {
305                entries.push(IntervalNode {
306                    high: coord,
307                    values: set,
308                });
309                self.size += 1;
310            }
311        }
312    }
313}
314
315pub struct BTreeEntry<'a, T: Clone + Eq + std::hash::Hash> {
316    tree: &'a mut IntervalTree<T>,
317    low: u32,
318    high: u32,
319}
320
321impl<'a, T: Clone + Eq + std::hash::Hash> BTreeEntry<'a, T> {
322    pub fn or_insert_with<F>(self, f: F) -> &'a mut HashSet<T>
323    where
324        F: FnOnce() -> HashSet<T>,
325    {
326        if self.tree.get_mut(self.low, self.high).is_none() {
327            let values = f();
328            let entries = self.tree.map.entry(self.low).or_default();
329            entries.push(IntervalNode {
330                high: self.high,
331                values,
332            });
333            self.tree.size += 1;
334        }
335        self.tree.get_mut(self.low, self.high).unwrap()
336    }
337}
338
339#[cfg(test)]
340mod tests {
341    use super::*;
342
343    #[test]
344    fn test_insert_and_query_point_interval() {
345        let mut tree = IntervalTree::new();
346        tree.insert(5, 5, 100);
347
348        let results = tree.query(5, 5);
349        assert_eq!(results.len(), 1);
350        assert_eq!(results[0].0, 5);
351        assert_eq!(results[0].1, 5);
352        assert!(results[0].2.contains(&100));
353    }
354
355    #[test]
356    fn test_insert_and_query_range() {
357        let mut tree = IntervalTree::new();
358        tree.insert(10, 20, 1);
359        tree.insert(15, 25, 2);
360        tree.insert(30, 40, 3);
361
362        // Query overlapping with first two intervals
363        let results = tree.query(12, 22);
364        assert_eq!(results.len(), 2);
365
366        // Query overlapping with only the third interval
367        let results = tree.query(35, 45);
368        assert_eq!(results.len(), 1);
369        assert!(results[0].2.contains(&3));
370    }
371
372    #[test]
373    fn point_interval_visit_does_not_scan_coordinate_prefixes() {
374        let mut tree = IntervalTree::new();
375        for coordinate in 0..10_000 {
376            tree.insert(coordinate, coordinate, coordinate);
377        }
378
379        let mut nodes = 0;
380        let mut values = Vec::new();
381        let result = tree.visit_point_intervals(9_999, 9_999, |entry| {
382            match entry {
383                None => nodes += 1,
384                Some(value) => values.push(*value),
385            }
386            ControlFlow::Continue(())
387        });
388
389        assert_eq!(result, ControlFlow::Continue(()));
390        assert_eq!(nodes, 1);
391        assert_eq!(values, vec![9_999]);
392    }
393
394    #[test]
395    fn point_interval_path_does_not_change_general_overlap_queries() {
396        let mut tree = IntervalTree::new();
397        tree.insert(1, 100, "range");
398        tree.insert(75, 75, "point");
399
400        let general = tree.query(75, 75);
401        assert_eq!(general.len(), 2);
402        assert!(general.iter().any(|entry| entry.2.contains("range")));
403        assert!(general.iter().any(|entry| entry.2.contains("point")));
404
405        let mut point_values = Vec::new();
406        let _ = tree.visit_point_intervals(75, 75, |entry| {
407            if let Some(value) = entry {
408                point_values.push(*value);
409            }
410            ControlFlow::Continue(())
411        });
412        assert_eq!(point_values, vec!["point"]);
413    }
414
415    #[test]
416    fn test_remove_value() {
417        let mut tree = IntervalTree::new();
418        tree.insert(5, 5, 100);
419        tree.insert(5, 5, 200);
420
421        assert_eq!(tree.query(5, 5).len(), 1);
422        assert_eq!(tree.query(5, 5)[0].2.len(), 2);
423
424        tree.remove(5, 5, &100);
425
426        let results = tree.query(5, 5);
427        assert_eq!(results.len(), 1);
428        assert_eq!(results[0].2.len(), 1);
429        assert!(results[0].2.contains(&200));
430    }
431
432    #[test]
433    fn test_entry_api() {
434        let mut tree: IntervalTree<i32> = IntervalTree::new();
435
436        tree.entry(10, 10).or_insert_with(HashSet::new).insert(42);
437
438        tree.entry(10, 10).or_insert_with(HashSet::new).insert(43);
439
440        let results = tree.query(10, 10);
441        assert_eq!(results.len(), 1);
442        assert_eq!(results[0].2.len(), 2);
443        assert!(results[0].2.contains(&42));
444        assert!(results[0].2.contains(&43));
445    }
446
447    #[test]
448    fn test_large_sparse_tree() {
449        let mut tree = IntervalTree::new();
450
451        // Simulate sparse spreadsheet
452        for i in (0..1_000_000).step_by(10000) {
453            tree.insert(i, i, i as i32);
454        }
455
456        assert_eq!(tree.len(), 100);
457
458        // Query for high rows
459        let results = tree.query(500_000, u32::MAX);
460        assert_eq!(results.len(), 50);
461    }
462
463    #[test]
464    fn test_entry_recursion_bug() {
465        let mut tree: IntervalTree<u32> = IntervalTree::new();
466
467        // The bug happens when we insert a value, then use entry()
468        // on a coordinate that would be a child of that value.
469        let count: u32 = 5000;
470        for i in 0..count {
471            tree.entry(i, i).or_insert_with(HashSet::new);
472        }
473
474        assert_eq!(tree.len(), count as usize);
475    }
476
477    #[test]
478    fn test_complex_overlaps() {
479        let mut tree = IntervalTree::new();
480        // Nested intervals
481        tree.insert(10, 100, "A");
482        tree.insert(20, 50, "B");
483        tree.insert(30, 40, "C");
484
485        // Partially overlapping
486        tree.insert(5, 15, "D");
487        tree.insert(95, 105, "E");
488
489        // Query for the very middle
490        let results = tree.query(35, 35);
491        assert_eq!(results.len(), 3); // Should hit A, B, and C
492
493        // Query for a range that only hits the "tail" of the large interval and the "head" of the end interval
494        let results = tree.query(98, 102);
495        assert_eq!(results.len(), 2); // Should hit A and E
496    }
497
498    #[test]
499    fn test_multiple_values_and_size() {
500        let mut tree = IntervalTree::new();
501
502        // Insert same interval twice with different values
503        tree.insert(10, 10, "val1");
504        tree.insert(10, 10, "val2");
505        assert_eq!(tree.len(), 1); // Size should only count unique intervals
506
507        // Insert same value twice
508        tree.insert(10, 10, "val1");
509        assert_eq!(tree.len(), 1);
510        let results = tree.query(10, 10);
511        assert_eq!(results[0].2.len(), 2); // HashSet handles the duplicate value "val1"
512    }
513
514    #[test]
515    fn test_remove_edge_cases() {
516        let mut tree = IntervalTree::new();
517        tree.insert(10, 20, "A");
518
519        // Try to remove a value that isn't there
520        let removed = tree.remove(10, 20, &"B");
521        assert!(!removed);
522        assert_eq!(tree.query(10, 20)[0].2.len(), 1);
523
524        // Try to remove from an interval that doesn't exist
525        let removed = tree.remove(99, 100, &"A");
526        assert!(!removed);
527    }
528
529    #[test]
530    fn test_bulk_build_consistency() {
531        let mut incremental_tree = IntervalTree::new();
532        let mut bulk_tree = IntervalTree::new();
533
534        let data: Vec<(u32, HashSet<&str>)> = vec![
535            (10, vec!["A", "B"].into_iter().collect()),
536            (20, vec!["C"].into_iter().collect()),
537            (5, vec!["D"].into_iter().collect()),
538        ];
539
540        // Build incrementally
541        for (coord, values) in &data {
542            for val in values {
543                incremental_tree.insert(*coord, *coord, *val);
544            }
545        }
546
547        // Build using bulk
548        bulk_tree.bulk_build_points(data.clone());
549
550        // Compare results
551        assert_eq!(incremental_tree.len(), bulk_tree.len());
552        assert_eq!(incremental_tree.query(0, 100), bulk_tree.query(0, 100));
553    }
554
555    #[test]
556    fn test_query_stack_safety() {
557        let mut tree = IntervalTree::new();
558        let count = 10_000;
559
560        // Create a deep right-leaning tree
561        for i in 0..count {
562            tree.insert(i, i, i);
563        }
564
565        // Query the very end of the tree
566        // If this causes a SIGABRT, it means query_node() must be made iterative
567        let results = tree.query(count - 1, count - 1);
568        assert_eq!(results.len(), 1);
569    }
570
571    #[test]
572    fn test_empty_and_boundaries() {
573        let mut tree: IntervalTree<i32> = IntervalTree::new();
574
575        assert!(tree.is_empty());
576        assert_eq!(tree.query(0, 100).len(), 0);
577        assert!(!tree.remove(0, 0, &1));
578
579        // Test a query that "misses" everything
580        tree.insert(50, 60, 1);
581        assert_eq!(tree.query(0, 49).len(), 0);
582        assert_eq!(tree.query(61, 100).len(), 0);
583    }
584
585    #[test]
586    fn test_multi_value_interval_size_tracking() {
587        let mut tree = IntervalTree::new();
588        let iv = (10, 20);
589
590        // 1. Insert two values for the same interval
591        // Destructure the tuple into low (iv.0) and high (iv.1)
592        tree.insert(iv.0, iv.1, "A");
593        tree.insert(iv.0, iv.1, "B");
594        assert_eq!(tree.len(), 1, "Should be 1 unique interval");
595
596        // 2. Remove first value - pass as reference &"A"
597        assert!(tree.remove(iv.0, iv.1, &"A"));
598        assert_eq!(
599            tree.len(),
600            1,
601            "Should still be 1 interval after partial removal"
602        );
603
604        // 3. Remove second value - size should now be 0
605        assert!(tree.remove(iv.0, iv.1, &"B"));
606        assert_eq!(tree.len(), 0, "Should be 0 after last value removed");
607    }
608}