Skip to main content

segment_tree/
lib.rs

1//! A generic segment tree.
2//!
3//! `T` is the type of the data in the array; `F` is the merge function
4//! (typically a sum, min, max, or xor closure).
5//!
6//! Both `query` and `update` run in `O(log n)`. The tree is stored in
7//! a flat `Vec<T>` of size `4 * n`, which is the standard bound for an
8//! implicit binary-tree representation.
9
10/// Generic segment tree.
11pub struct SegmentTree<T, F> {
12    size: usize,
13    tree: Vec<T>,
14    default: T,
15    merge: F,
16}
17
18impl<T, F> SegmentTree<T, F>
19where
20    T: Clone + Copy,
21    F: Fn(T, T) -> T,
22{
23    /// Build a segment tree over `arr` using `default` as the identity
24    /// for `merge` and `merge` itself as the associative combining op.
25    pub fn new(arr: &[T], default: T, merge: F) -> Self {
26        let size = arr.len();
27        #[allow(unused_mut)]
28        let mut tree = vec![default; 4 * size];
29        let mut seg_tree = SegmentTree {
30            size,
31            tree,
32            default,
33            merge,
34        };
35        seg_tree.build(arr, 0, 0, size - 1);
36        seg_tree
37    }
38
39    fn build(&mut self, arr: &[T], node: usize, start: usize, end: usize) {
40        if start == end {
41            self.tree[node] = arr[start];
42            return;
43        }
44        let mid = (start + end) / 2;
45        self.build(arr, 2 * node + 1, start, mid);
46        self.build(arr, 2 * node + 2, mid + 1, end);
47        self.tree[node] = (self.merge)(self.tree[2 * node + 1], self.tree[2 * node + 2]);
48    }
49
50    /// Query `merge` over the closed range `[l, r]`.
51    pub fn query(&self, l: usize, r: usize) -> T {
52        self.query_recursive(0, 0, self.size - 1, l, r)
53    }
54
55    fn query_recursive(&self, node: usize, start: usize, end: usize, l: usize, r: usize) -> T {
56        if r < start || end < l {
57            return self.default;
58        }
59        if l <= start && end <= r {
60            return self.tree[node];
61        }
62        let mid = (start + end) / 2;
63        let p1 = self.query_recursive(2 * node + 1, start, mid, l, r);
64        let p2 = self.query_recursive(2 * node + 2, mid + 1, end, l, r);
65        (self.merge)(p1, p2)
66    }
67
68    /// Point update: set `arr[idx] = val` and recompute affected nodes.
69    pub fn update(&mut self, idx: usize, val: T) {
70        self.update_recursive(0, 0, self.size - 1, idx, val);
71    }
72
73    fn update_recursive(&mut self, node: usize, start: usize, end: usize, idx: usize, val: T) {
74        if start == end {
75            self.tree[node] = val;
76            return;
77        }
78        let mid = (start + end) / 2;
79        if start <= idx && idx <= mid {
80            self.update_recursive(2 * node + 1, start, mid, idx, val);
81        } else {
82            self.update_recursive(2 * node + 2, mid + 1, end, idx, val);
83        }
84        self.tree[node] = (self.merge)(self.tree[2 * node + 1], self.tree[2 * node + 2]);
85    }
86}