Skip to main content

formualizer_eval/engine/
sheet_index.rs

1use super::interval_tree::IntervalTree;
2use super::vertex::VertexId;
3use formualizer_common::Coord as AbsCoord;
4use std::collections::HashSet;
5use std::ops::ControlFlow;
6#[cfg(test)]
7use std::sync::atomic::{AtomicUsize, Ordering};
8
9#[cfg(test)]
10#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
11pub(crate) struct SheetIndexQueryStats {
12    pub coordinate_nodes_visited: usize,
13    pub values_visited: usize,
14}
15
16/// Sheet-level sparse index for efficient range queries on vertex positions.
17///
18/// ## Why SheetIndex with interval trees?
19///
20/// While `cell_to_vertex: HashMap<CellRef, VertexId>` provides O(1) exact lookups,
21/// structural operations (insert/delete rows/columns) need to find ALL vertices
22/// in a given range, which would require O(n) full scans of the hash map.
23///
24/// ### Performance comparison:
25///
26/// | Operation | Hash map only | With SheetIndex |
27/// |-----------|---------------|-----------------|
28/// | Insert 100 rows at row 20,000 | O(total cells) | O(log n + k)* |
29/// | Delete columns B:D | O(total cells) | O(log n + k) |
30/// | Viewport query (visible cells) | O(total cells) | O(log n + k) |
31///
32/// *where n = number of indexed vertices, k = vertices actually affected
33///
34/// ### Memory efficiency:
35///
36/// - Each interval is just 2×u32 + Vec<VertexId> pointer
37/// - Spreadsheets are extremely sparse (1M row sheet typically has <10K cells)
38/// - Point intervals (single cells) are the common case
39/// - Trees stay small and cache-friendly
40///
41/// ### Future benefits:
42///
43/// 1. **Virtual scrolling** - fetch viewport cells in microseconds
44/// 2. **Lazy evaluation** - mark row blocks dirty without scanning
45/// 3. **Concurrent reads** - trees are read-mostly, perfect for RwLock
46/// 4. **Minimal undo/redo** - know exactly which vertices were touched
47#[derive(Debug, Default)]
48pub struct SheetIndex {
49    memberships: HashSet<VertexId>,
50    /// Row interval tree: maps row ranges → vertices in those rows
51    /// For a cell at (r,c), we store the point interval [r,r] → VertexId
52    row_tree: IntervalTree<VertexId>,
53
54    /// Column interval tree: maps column ranges → vertices in those columns  
55    /// For a cell at (r,c), we store the point interval [c,c] → VertexId
56    col_tree: IntervalTree<VertexId>,
57
58    #[cfg(test)]
59    query_coordinate_nodes_visited: AtomicUsize,
60    #[cfg(test)]
61    query_values_visited: AtomicUsize,
62}
63
64impl SheetIndex {
65    /// Create a new empty sheet index
66    pub fn new() -> Self {
67        Self {
68            memberships: HashSet::new(),
69            row_tree: IntervalTree::new(),
70            col_tree: IntervalTree::new(),
71            #[cfg(test)]
72            query_coordinate_nodes_visited: AtomicUsize::new(0),
73            #[cfg(test)]
74            query_values_visited: AtomicUsize::new(0),
75        }
76    }
77
78    /// Fast path build from sorted coordinates. Assumes items are row-major sorted.
79    pub fn build_from_sorted(&mut self, items: &[(AbsCoord, VertexId)]) {
80        self.add_vertices_batch(items);
81    }
82
83    /// Add a vertex at the given coordinate to the index.
84    ///
85    /// ## Complexity
86    /// O(log n) where n is the number of vertices in the index
87    pub fn add_vertex(&mut self, coord: AbsCoord, vertex_id: VertexId) {
88        let row = coord.row();
89        let col = coord.col();
90
91        if !self.memberships.insert(vertex_id) {
92            return;
93        }
94
95        // Add to row tree - point interval [row, row]
96        self.row_tree
97            .entry(row, row)
98            .or_insert_with(HashSet::new)
99            .insert(vertex_id);
100
101        // Add to column tree - point interval [col, col]
102        self.col_tree
103            .entry(col, col)
104            .or_insert_with(HashSet::new)
105            .insert(vertex_id);
106    }
107
108    /// Add many vertices in a single pass. Assumes coords belong to same sheet index.
109    pub fn add_vertices_batch(&mut self, items: &[(AbsCoord, VertexId)]) {
110        if items.is_empty() {
111            return;
112        }
113        // If trees are empty we can bulk build from sorted points in O(n log n) with better constants.
114        if self.row_tree.is_empty() && self.col_tree.is_empty() {
115            // Build row points
116            let mut row_items: Vec<(u32, HashSet<VertexId>)> = Vec::with_capacity(items.len());
117            let mut col_items: Vec<(u32, HashSet<VertexId>)> = Vec::with_capacity(items.len());
118            // Use temp hash maps for merging duplicates
119            use rustc_hash::FxHashMap;
120            let mut row_map: FxHashMap<u32, HashSet<VertexId>> = FxHashMap::default();
121            let mut col_map: FxHashMap<u32, HashSet<VertexId>> = FxHashMap::default();
122            for (coord, vid) in items {
123                if !self.memberships.insert(*vid) {
124                    continue;
125                }
126                row_map.entry(coord.row()).or_default().insert(*vid);
127                col_map.entry(coord.col()).or_default().insert(*vid);
128            }
129            row_items.reserve(row_map.len());
130            for (k, v) in row_map.into_iter() {
131                row_items.push((k, v));
132            }
133            col_items.reserve(col_map.len());
134            for (k, v) in col_map.into_iter() {
135                col_items.push((k, v));
136            }
137            self.row_tree.bulk_build_points(row_items);
138            self.col_tree.bulk_build_points(col_items);
139            return;
140        }
141        // Fallback: incremental for already populated index
142        for (coord, vid) in items {
143            self.add_vertex(*coord, *vid);
144        }
145    }
146
147    /// Remove a vertex from the index.
148    ///
149    /// ## Complexity
150    /// O(log n) where n is the number of vertices in the index
151    pub fn remove_vertex(&mut self, coord: AbsCoord, vertex_id: VertexId) {
152        let row = coord.row();
153        let col = coord.col();
154
155        if !self.memberships.remove(&vertex_id) {
156            return;
157        }
158
159        self.row_tree.remove(row, row, &vertex_id);
160        self.col_tree.remove(col, col, &vertex_id);
161    }
162
163    /// Update a vertex's position in the index (move operation).
164    ///
165    /// ## Complexity
166    /// O(log n) for removal + O(log n) for insertion = O(log n)
167    pub fn update_vertex(&mut self, old_coord: AbsCoord, new_coord: AbsCoord, vertex_id: VertexId) {
168        self.remove_vertex(old_coord, vertex_id);
169        self.add_vertex(new_coord, vertex_id);
170    }
171
172    fn record_coordinate_visits(&self, count: usize) {
173        #[cfg(test)]
174        self.query_coordinate_nodes_visited
175            .fetch_add(count, Ordering::Relaxed);
176        #[cfg(not(test))]
177        let _ = count;
178    }
179
180    fn record_value_visit(&self) {
181        #[cfg(test)]
182        self.query_values_visited.fetch_add(1, Ordering::Relaxed);
183    }
184
185    fn visit_axis_range(
186        &self,
187        tree: &IntervalTree<VertexId>,
188        start: u32,
189        end: u32,
190        mut visitor: impl FnMut(VertexId),
191    ) {
192        let _ = tree.visit_point_intervals(start, end, |entry| {
193            match entry {
194                None => self.record_coordinate_visits(1),
195                Some(vertex) => {
196                    self.record_value_visit();
197                    visitor(*vertex);
198                }
199            }
200            ControlFlow::Continue(())
201        });
202    }
203
204    fn axis_range_value_count(&self, tree: &IntervalTree<VertexId>, start: u32, end: u32) -> usize {
205        let (nodes, values) = tree.point_interval_stats(start, end);
206        self.record_coordinate_visits(nodes);
207        values
208    }
209
210    fn collect_axis_range(
211        &self,
212        tree: &IntervalTree<VertexId>,
213        start: u32,
214        end: u32,
215    ) -> HashSet<VertexId> {
216        let mut result = HashSet::new();
217        self.visit_axis_range(tree, start, end, |vertex| {
218            result.insert(vertex);
219        });
220        result
221    }
222
223    /// Query all vertices in the given row range.
224    ///
225    /// ## Complexity
226    /// O(log n + k) where k is the number of vertices in the range
227    pub fn vertices_in_row_range(&self, start: u32, end: u32) -> Vec<VertexId> {
228        self.collect_axis_range(&self.row_tree, start, end)
229            .into_iter()
230            .collect()
231    }
232
233    /// Query all vertices in the given column range.
234    ///
235    /// ## Complexity
236    /// O(log n + k) where k is the number of vertices in the range
237    pub fn vertices_in_col_range(&self, start: u32, end: u32) -> Vec<VertexId> {
238        self.collect_axis_range(&self.col_tree, start, end)
239            .into_iter()
240            .collect()
241    }
242
243    /// Query all vertices in a rectangular range.
244    ///
245    /// Sheet indexes contain point intervals only, so exact-cell queries use two
246    /// direct B-tree lookups. Wider rectangles materialize only the cheaper axis
247    /// set and stream the other axis while intersecting it.
248    pub fn vertices_in_rect(
249        &self,
250        start_row: u32,
251        end_row: u32,
252        start_col: u32,
253        end_col: u32,
254    ) -> Vec<VertexId> {
255        if start_row > end_row || start_col > end_col {
256            return Vec::new();
257        }
258
259        if start_row == end_row && start_col == end_col {
260            self.record_coordinate_visits(2);
261            let Some(row_vertices) = self.row_tree.point_values(start_row) else {
262                return Vec::new();
263            };
264            let Some(col_vertices) = self.col_tree.point_values(start_col) else {
265                return Vec::new();
266            };
267            let (candidates, membership) = if row_vertices.len() <= col_vertices.len() {
268                (row_vertices, col_vertices)
269            } else {
270                (col_vertices, row_vertices)
271            };
272            return candidates
273                .iter()
274                .filter_map(|vertex| {
275                    self.record_value_visit();
276                    membership.contains(vertex).then_some(*vertex)
277                })
278                .collect();
279        }
280
281        let row_count = self.axis_range_value_count(&self.row_tree, start_row, end_row);
282        let col_count = self.axis_range_value_count(&self.col_tree, start_col, end_col);
283        let (candidates, other_tree, other_start, other_end) = if row_count <= col_count {
284            (
285                self.collect_axis_range(&self.row_tree, start_row, end_row),
286                &self.col_tree,
287                start_col,
288                end_col,
289            )
290        } else {
291            (
292                self.collect_axis_range(&self.col_tree, start_col, end_col),
293                &self.row_tree,
294                start_row,
295                end_row,
296            )
297        };
298        let mut result = Vec::with_capacity(candidates.len().min(row_count).min(col_count));
299        self.visit_axis_range(other_tree, other_start, other_end, |vertex| {
300            if candidates.contains(&vertex) {
301                result.push(vertex);
302            }
303        });
304        result
305    }
306
307    #[cfg(test)]
308    pub(crate) fn reset_query_stats(&self) {
309        self.query_coordinate_nodes_visited
310            .store(0, Ordering::Relaxed);
311        self.query_values_visited.store(0, Ordering::Relaxed);
312    }
313
314    #[cfg(test)]
315    pub(crate) fn query_stats(&self) -> SheetIndexQueryStats {
316        SheetIndexQueryStats {
317            coordinate_nodes_visited: self.query_coordinate_nodes_visited.load(Ordering::Relaxed),
318            values_visited: self.query_values_visited.load(Ordering::Relaxed),
319        }
320    }
321
322    pub fn len(&self) -> usize {
323        self.memberships.len()
324    }
325
326    /// Check if the index is empty.
327    pub fn is_empty(&self) -> bool {
328        self.memberships.is_empty()
329    }
330
331    /// Clear all entries from the index.
332    pub fn clear(&mut self) {
333        self.memberships.clear();
334        self.row_tree = IntervalTree::new();
335        self.col_tree = IntervalTree::new();
336    }
337}
338
339#[cfg(test)]
340mod tests {
341    use super::*;
342
343    #[test]
344    fn test_add_and_query_single_vertex() {
345        let mut index = SheetIndex::new();
346        let coord = AbsCoord::new(5, 10);
347        let vertex_id = VertexId(1024);
348
349        index.add_vertex(coord, vertex_id);
350
351        // Query exact row
352        let row_results = index.vertices_in_row_range(5, 5);
353        assert_eq!(row_results, vec![vertex_id]);
354
355        // Query exact column
356        let col_results = index.vertices_in_col_range(10, 10);
357        assert_eq!(col_results, vec![vertex_id]);
358
359        // Query range containing the vertex
360        let row_results = index.vertices_in_row_range(3, 7);
361        assert_eq!(row_results, vec![vertex_id]);
362    }
363
364    #[test]
365    fn vertex_count_is_unique_and_consistent_across_incremental_and_batch_builds() {
366        let vertex = VertexId(1024);
367        let items = [(AbsCoord::new(1, 1), vertex), (AbsCoord::new(1, 1), vertex)];
368        let mut incremental = SheetIndex::new();
369        for (coord, vertex) in items {
370            incremental.add_vertex(coord, vertex);
371        }
372        let mut batch = SheetIndex::new();
373        batch.add_vertices_batch(&items);
374        assert_eq!(incremental.len(), 1);
375        assert_eq!(batch.len(), incremental.len());
376    }
377
378    #[test]
379    fn test_remove_vertex() {
380        let mut index = SheetIndex::new();
381        let coord = AbsCoord::new(5, 10);
382        let vertex_id = VertexId(1024);
383
384        index.add_vertex(coord, vertex_id);
385        assert_eq!(index.len(), 1);
386
387        index.remove_vertex(coord, vertex_id);
388        assert_eq!(index.len(), 0);
389
390        // Should return empty after removal
391        let row_results = index.vertices_in_row_range(5, 5);
392        assert!(row_results.is_empty());
393    }
394
395    #[test]
396    fn test_update_vertex_position() {
397        let mut index = SheetIndex::new();
398        let old_coord = AbsCoord::new(5, 10);
399        let new_coord = AbsCoord::new(15, 20);
400        let vertex_id = VertexId(1024);
401
402        index.add_vertex(old_coord, vertex_id);
403        index.update_vertex(old_coord, new_coord, vertex_id);
404
405        // Should not be at old position
406        let old_row_results = index.vertices_in_row_range(5, 5);
407        assert!(old_row_results.is_empty());
408
409        // Should be at new position
410        let new_row_results = index.vertices_in_row_range(15, 15);
411        assert_eq!(new_row_results, vec![vertex_id]);
412
413        let new_col_results = index.vertices_in_col_range(20, 20);
414        assert_eq!(new_col_results, vec![vertex_id]);
415    }
416
417    #[test]
418    fn test_range_queries() {
419        let mut index = SheetIndex::new();
420
421        // Add vertices in a pattern
422        for row in 0..10 {
423            for col in 0..5 {
424                let coord = AbsCoord::new(row, col);
425                let vertex_id = VertexId(1024 + row * 5 + col);
426                index.add_vertex(coord, vertex_id);
427            }
428        }
429
430        // Query rows 3-5 (should get 3 rows × 5 cols = 15 vertices)
431        let row_results = index.vertices_in_row_range(3, 5);
432        assert_eq!(row_results.len(), 15);
433
434        // Query columns 1-2 (should get 10 rows × 2 cols = 20 vertices)
435        let col_results = index.vertices_in_col_range(1, 2);
436        assert_eq!(col_results.len(), 20);
437
438        // Query rectangle (rows 3-5, cols 1-2) should get 3 × 2 = 6 vertices
439        let rect_results = index.vertices_in_rect(3, 5, 1, 2);
440        assert_eq!(rect_results.len(), 6);
441    }
442
443    #[test]
444    fn test_sparse_sheet_efficiency() {
445        let mut index = SheetIndex::new();
446
447        // Simulate sparse sheet - only a few cells in a million-row range
448        index.add_vertex(AbsCoord::new(100, 5), VertexId(1024));
449        index.add_vertex(AbsCoord::new(50_000, 10), VertexId(1025));
450        index.add_vertex(AbsCoord::new(100_000, 15), VertexId(1026));
451        index.add_vertex(AbsCoord::new(500_000, 20), VertexId(1027));
452        index.add_vertex(AbsCoord::new(999_999, 25), VertexId(1028));
453
454        assert_eq!(index.len(), 5);
455
456        // Query for rows >= 100_000 (should find 3 vertices efficiently)
457        let high_rows = index.vertices_in_row_range(100_000, u32::MAX);
458        assert_eq!(high_rows.len(), 3);
459
460        // Query for specific column range
461        let col_range = index.vertices_in_col_range(10, 20);
462        assert_eq!(col_range.len(), 3); // columns 10, 15, 20
463    }
464
465    #[test]
466    fn test_shift_operation_query() {
467        let mut index = SheetIndex::new();
468
469        // Setup: cells at rows 10, 20, 30, 40, 50
470        for row in [10, 20, 30, 40, 50] {
471            index.add_vertex(AbsCoord::new(row, 0), VertexId(1024 + row));
472        }
473
474        // Simulate "insert 5 rows at row 25" - need to find all vertices with row >= 25
475        let vertices_to_shift = index.vertices_in_row_range(25, u32::MAX);
476        assert_eq!(vertices_to_shift.len(), 3); // rows 30, 40, 50
477
478        // Simulate "delete columns B:D" - need to find all vertices in columns 1-3
479        for col in 1..=3 {
480            index.add_vertex(AbsCoord::new(5, col), VertexId(2000 + col));
481        }
482
483        let vertices_to_delete = index.vertices_in_col_range(1, 3);
484        assert_eq!(vertices_to_delete.len(), 3);
485    }
486
487    #[test]
488    fn test_viewport_query() {
489        let mut index = SheetIndex::new();
490
491        // Simulate a spreadsheet with scattered data
492        for row in (0..10000).step_by(100) {
493            for col in 0..10 {
494                index.add_vertex(AbsCoord::new(row, col), VertexId(row * 10 + col));
495            }
496        }
497
498        // Query viewport: rows 500-1500, columns 2-7
499        let viewport = index.vertices_in_rect(500, 1500, 2, 7);
500
501        // Should find 11 rows (500, 600, ..., 1500) × 6 columns (2-7) = 66 vertices
502        assert_eq!(viewport.len(), 66);
503    }
504
505    #[test]
506    fn exact_cell_query_visits_only_exact_coordinate_buckets() {
507        let mut index = SheetIndex::new();
508        for row in 0..10_000 {
509            index.add_vertex(AbsCoord::new(row, 7), VertexId(1024 + row));
510        }
511
512        index.reset_query_stats();
513        assert_eq!(
514            index.vertices_in_rect(9_999, 9_999, 7, 7),
515            vec![VertexId(11_023)]
516        );
517        assert_eq!(
518            index.query_stats(),
519            SheetIndexQueryStats {
520                coordinate_nodes_visited: 2,
521                values_visited: 1,
522            }
523        );
524    }
525
526    fn sorted(mut vertices: Vec<VertexId>) -> Vec<VertexId> {
527        vertices.sort_unstable();
528        vertices
529    }
530
531    fn assert_query_parity(
532        index: &SheetIndex,
533        model: &[(AbsCoord, VertexId)],
534        start_row: u32,
535        end_row: u32,
536        start_col: u32,
537        end_col: u32,
538    ) {
539        let naive_rect = model
540            .iter()
541            .filter_map(|(coord, vertex)| {
542                (coord.row() >= start_row
543                    && coord.row() <= end_row
544                    && coord.col() >= start_col
545                    && coord.col() <= end_col)
546                    .then_some(*vertex)
547            })
548            .collect::<Vec<_>>();
549        let naive_rows = model
550            .iter()
551            .filter_map(|(coord, vertex)| {
552                (coord.row() >= start_row && coord.row() <= end_row).then_some(*vertex)
553            })
554            .collect::<Vec<_>>();
555        let naive_cols = model
556            .iter()
557            .filter_map(|(coord, vertex)| {
558                (coord.col() >= start_col && coord.col() <= end_col).then_some(*vertex)
559            })
560            .collect::<Vec<_>>();
561        assert_eq!(
562            sorted(index.vertices_in_rect(start_row, end_row, start_col, end_col)),
563            sorted(naive_rect)
564        );
565        assert_eq!(
566            sorted(index.vertices_in_row_range(start_row, end_row)),
567            sorted(naive_rows)
568        );
569        assert_eq!(
570            sorted(index.vertices_in_col_range(start_col, end_col)),
571            sorted(naive_cols)
572        );
573    }
574
575    #[test]
576    fn point_range_rectangle_sparse_move_remove_and_randomized_queries_match_naive_filtering() {
577        let mut state = 0x5eed_cafe_u64;
578        let mut next = || {
579            state = state
580                .wrapping_mul(6_364_136_223_846_793_005)
581                .wrapping_add(1);
582            (state >> 32) as u32
583        };
584        let mut model = (0..300)
585            .map(|offset| {
586                let row = if offset < 5 {
587                    offset * 200_000
588                } else {
589                    next() % 2_000
590                };
591                let col = if offset < 5 { offset * 20 } else { next() % 80 };
592                (AbsCoord::new(row, col), VertexId(1024 + offset))
593            })
594            .collect::<Vec<_>>();
595
596        let mut incremental = SheetIndex::new();
597        for &(coord, vertex) in &model {
598            incremental.add_vertex(coord, vertex);
599        }
600        let mut bulk_items = model.clone();
601        bulk_items.sort_unstable_by_key(|(coord, _)| (coord.row(), coord.col()));
602        let mut bulk = SheetIndex::new();
603        bulk.build_from_sorted(&bulk_items);
604
605        for index in [&incremental, &bulk] {
606            assert_query_parity(index, &model, 1_500, 1_500, 40, 40);
607            assert_query_parity(index, &model, 500, 1_500, 0, 79);
608            assert_query_parity(index, &model, 0, u32::MAX, 20, 40);
609            assert_query_parity(index, &model, 0, 800_000, 0, 80);
610        }
611
612        for (offset, entry) in model.iter_mut().take(40).enumerate() {
613            let (old_coord, vertex) = *entry;
614            let new_coord = AbsCoord::new(3_000 + offset as u32, 100 + offset as u32 % 7);
615            incremental.update_vertex(old_coord, new_coord, vertex);
616            bulk.update_vertex(old_coord, new_coord, vertex);
617            entry.0 = new_coord;
618        }
619        for _ in 0..30 {
620            let index = (next() as usize) % model.len();
621            let (coord, vertex) = model.swap_remove(index);
622            incremental.remove_vertex(coord, vertex);
623            bulk.remove_vertex(coord, vertex);
624        }
625        let point = model[0].0;
626        for index in [&incremental, &bulk] {
627            assert_query_parity(
628                index,
629                &model,
630                point.row(),
631                point.row(),
632                point.col(),
633                point.col(),
634            );
635        }
636
637        for _ in 0..200 {
638            let row_a = next() % 4_000;
639            let row_b = next() % 4_000;
640            let col_a = next() % 120;
641            let col_b = next() % 120;
642            let (start_row, end_row) = (row_a.min(row_b), row_a.max(row_b));
643            let (start_col, end_col) = (col_a.min(col_b), col_a.max(col_b));
644            assert_query_parity(&incremental, &model, start_row, end_row, start_col, end_col);
645            assert_query_parity(&bulk, &model, start_row, end_row, start_col, end_col);
646        }
647    }
648}