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#[derive(Debug, Default)]
48pub struct SheetIndex {
49 memberships: HashSet<VertexId>,
50 row_tree: IntervalTree<VertexId>,
53
54 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 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 pub fn build_from_sorted(&mut self, items: &[(AbsCoord, VertexId)]) {
80 self.add_vertices_batch(items);
81 }
82
83 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 self.row_tree
97 .entry(row, row)
98 .or_insert_with(HashSet::new)
99 .insert(vertex_id);
100
101 self.col_tree
103 .entry(col, col)
104 .or_insert_with(HashSet::new)
105 .insert(vertex_id);
106 }
107
108 pub fn add_vertices_batch(&mut self, items: &[(AbsCoord, VertexId)]) {
110 if items.is_empty() {
111 return;
112 }
113 if self.row_tree.is_empty() && self.col_tree.is_empty() {
115 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 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 for (coord, vid) in items {
143 self.add_vertex(*coord, *vid);
144 }
145 }
146
147 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 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 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 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 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 pub fn is_empty(&self) -> bool {
328 self.memberships.is_empty()
329 }
330
331 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 let row_results = index.vertices_in_row_range(5, 5);
353 assert_eq!(row_results, vec![vertex_id]);
354
355 let col_results = index.vertices_in_col_range(10, 10);
357 assert_eq!(col_results, vec![vertex_id]);
358
359 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 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 let old_row_results = index.vertices_in_row_range(5, 5);
407 assert!(old_row_results.is_empty());
408
409 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 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 let row_results = index.vertices_in_row_range(3, 5);
432 assert_eq!(row_results.len(), 15);
433
434 let col_results = index.vertices_in_col_range(1, 2);
436 assert_eq!(col_results.len(), 20);
437
438 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 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 let high_rows = index.vertices_in_row_range(100_000, u32::MAX);
458 assert_eq!(high_rows.len(), 3);
459
460 let col_range = index.vertices_in_col_range(10, 20);
462 assert_eq!(col_range.len(), 3); }
464
465 #[test]
466 fn test_shift_operation_query() {
467 let mut index = SheetIndex::new();
468
469 for row in [10, 20, 30, 40, 50] {
471 index.add_vertex(AbsCoord::new(row, 0), VertexId(1024 + row));
472 }
473
474 let vertices_to_shift = index.vertices_in_row_range(25, u32::MAX);
476 assert_eq!(vertices_to_shift.len(), 3); 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 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 let viewport = index.vertices_in_rect(500, 1500, 2, 7);
500
501 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}