1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
use std::cmp::{max, min};
use std::ops::Add;
use crate::node::*;
pub struct Tree<T: Clone + Add> {
pub n: usize,
pub root: NodePtr<T>,
pub combine: fn(a: T, b: T) -> T,
pub nil: T,
}
impl<T: Clone + Add<Output = T>> Tree<T> {
/// Returns a segment tree, that can query the combine() function on segment.
///
/// Arguments:
///
/// - `x` - initial Vec<T>
/// - `combine` - function to query
/// - `nil` - nil element; for sum it is `0`, max - `inf`, min - `-inf`, gcd - `0`
/// ```
/// use segtree::{sum, tree::Tree};
///
/// let a = vec![1, 2, 3, 4, 5, 6];
/// let tr = Tree::new(&a, segtree::sum, 0);
/// ```
pub fn new(x: &Vec<T>, combine: fn(a: T, b: T) -> T, nil: T) -> Self {
let mut tr = Tree {
n: x.len(),
root: None,
combine,
nil,
};
tr.root = tr.build(x, 0, tr.n);
return tr;
}
/// Builds a segment tree.
///
/// Complexity: O(n)
pub fn build(&self, x: &Vec<T>, lv: usize, rv: usize) -> NodePtr<T> {
if rv - lv == 1 {
return Node::new(x[lv].clone());
}
let mid = (lv + rv) / 2;
let left = self.build(x, lv, mid);
let right = self.build(x, mid, rv);
return Node::merge(
left.clone(),
right.clone(),
(self.combine)(left.unwrap().val, right.unwrap().val),
);
}
fn __get(&self, v: &NodePtr<T>, lv: usize, rv: usize, l: usize, r: usize) -> T {
if r <= lv || rv <= l {
return self.nil.clone();
}
if l <= lv && rv <= r {
return v.clone().unwrap().val;
}
let mid = (lv + rv) / 2;
return (self.combine)(
self.__get(&v.clone().unwrap().l, lv, mid, l, min(r, mid)),
self.__get(&v.clone().unwrap().r, mid, rv, max(l, mid), r),
);
}
/// Queries self.combine() on segment [l, r). So, r will not be included in the segment.
///
/// Complexity: O(log n)
/// ```
/// use segtree::{sum, tree::Tree};
///
/// let a = vec![1, 2, 3, 4, 5];
/// let tr = Tree::new(&a, sum, 0);
/// assert_eq!(tr.get(0, 3), 6); // a[0] + a[1] + a[2] = 1 + 2 + 3 = 6
/// ```
pub fn get(&self, l: usize, r: usize) -> T {
return self.__get(&self.root, 0, self.n, l, r);
}
fn __update(&self, v: &NodePtr<T>, lv: usize, rv: usize, ind: usize, val: T) -> NodePtr<T> {
if rv - lv == 1 {
return Node::new(val);
}
let mid = (lv + rv) / 2;
if ind < mid {
let left = self.__update(&v.clone().unwrap().l, lv, mid, ind, val);
let right = v.clone().unwrap().r;
return Node::merge(
left.clone(),
right.clone(),
(self.combine)(left.unwrap().val, right.unwrap().val),
);
} else {
let left = v.clone().unwrap().l;
let right = self.__update(&v.clone().unwrap().r, mid, rv, ind, val);
return Node::merge(
left.clone(),
right.clone(),
(self.combine)(left.unwrap().val, right.unwrap().val),
);
}
}
/// Queries an update on the given index. *Doesn't update initial Vec*
///
/// Complexity: O(log n)
/// ```
/// use segtree::{sum, tree::Tree};
///
/// let a = vec![1, 2, 3];
/// let mut tr = Tree::new(&a, sum, 0);
/// assert_eq!(tr.get(1, 2), 2); // a[1] = 2
/// tr.update(1, 100);
/// assert_eq!(tr.get(1, 2), 100); // a[1] is now equals to 100
/// ```
pub fn update(&mut self, ind: usize, val: T) {
self.root = self.__update(&self.root, 0, self.n, ind, val);
}
}