1use std::collections::{BTreeMap, HashSet};
2use std::ops::ControlFlow;
3
4#[derive(Debug, Clone)]
23struct IntervalNode<T: Clone + Eq + std::hash::Hash> {
24 high: u32,
25 values: HashSet<T>,
26}
27
28#[derive(Debug, Clone)]
30pub struct IntervalTree<T: Clone + Eq + std::hash::Hash> {
31 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 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 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 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 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 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 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 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 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 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 pub fn bulk_build_points(&mut self, mut items: Vec<(u32, HashSet<T>)>) {
280 if !self.is_empty() {
281 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 items.sort_by_key(|(k, _)| *k);
296
297 for (coord, set) in items {
299 let entries = self.map.entry(coord).or_default();
300
301 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 let results = tree.query(12, 22);
364 assert_eq!(results.len(), 2);
365
366 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 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 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 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 tree.insert(10, 100, "A");
482 tree.insert(20, 50, "B");
483 tree.insert(30, 40, "C");
484
485 tree.insert(5, 15, "D");
487 tree.insert(95, 105, "E");
488
489 let results = tree.query(35, 35);
491 assert_eq!(results.len(), 3); let results = tree.query(98, 102);
495 assert_eq!(results.len(), 2); }
497
498 #[test]
499 fn test_multiple_values_and_size() {
500 let mut tree = IntervalTree::new();
501
502 tree.insert(10, 10, "val1");
504 tree.insert(10, 10, "val2");
505 assert_eq!(tree.len(), 1); 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); }
513
514 #[test]
515 fn test_remove_edge_cases() {
516 let mut tree = IntervalTree::new();
517 tree.insert(10, 20, "A");
518
519 let removed = tree.remove(10, 20, &"B");
521 assert!(!removed);
522 assert_eq!(tree.query(10, 20)[0].2.len(), 1);
523
524 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 for (coord, values) in &data {
542 for val in values {
543 incremental_tree.insert(*coord, *coord, *val);
544 }
545 }
546
547 bulk_tree.bulk_build_points(data.clone());
549
550 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 for i in 0..count {
562 tree.insert(i, i, i);
563 }
564
565 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 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 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 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 assert!(tree.remove(iv.0, iv.1, &"B"));
606 assert_eq!(tree.len(), 0, "Should be 0 after last value removed");
607 }
608}