segtree 0.1.0

Segment tree implementation in rust
Documentation
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);
    }
}