use itertools::Itertools;
use vers_vecs::FastRmq;
use crate::iter::node_iter::EulerWalk;
use crate::tree::simple_rtree::TreeNodeID;
const NO_FIRST_APPEARANCE: u32 = u32::MAX;
pub struct LcaOracle<'t, Tree: EulerWalk> {
tree: &'t Tree,
euler: Vec<TreeNodeID<Tree>>,
fai: Vec<u32>,
da: Vec<usize>,
rmq: FastRmq,
}
impl<'t, Tree: EulerWalk> LcaOracle<'t, Tree> {
pub(crate) fn build(tree: &'t Tree) -> Self {
let euler: Vec<TreeNodeID<Tree>> = tree.euler_walk_ids(tree.get_root_id()).collect_vec();
let max_id: usize = tree
.get_node_ids()
.map(Into::<usize>::into)
.max()
.expect("a tree always has at least a root node");
let mut fai = vec![NO_FIRST_APPEARANCE; max_id + 1];
for (pos, node_id) in euler.iter().enumerate() {
let idx: usize = (*node_id).into();
if fai[idx] == NO_FIRST_APPEARANCE {
fai[idx] = pos as u32;
}
}
let root_id = tree.get_root_id();
let mut node_depth = vec![0usize; max_id + 1];
let mut stack = vec![(root_id, 0usize)];
while let Some((node_id, depth)) = stack.pop() {
node_depth[Into::<usize>::into(node_id)] = depth;
for child_id in tree.get_node_children_ids(node_id) {
stack.push((child_id, depth + 1));
}
}
let da: Vec<usize> = euler
.iter()
.map(|x| node_depth[Into::<usize>::into(*x)])
.collect();
let rmq = FastRmq::from_vec(da.iter().map(|x| *x as u64).collect_vec());
LcaOracle {
tree,
euler,
fai,
da,
rmq,
}
}
pub fn get_lca_id(&self, node_id_vec: &[TreeNodeID<Tree>]) -> TreeNodeID<Tree> {
if node_id_vec.len() == 1 {
return node_id_vec[0];
}
let (min_pos, max_pos) = node_id_vec
.iter()
.map(|x| self.get_fa_index(*x))
.fold((usize::MAX, usize::MIN), |(lo, hi), pos| {
(lo.min(pos), hi.max(pos))
});
self.euler[self.rmq.range_min(min_pos, max_pos)]
}
pub fn get_lca(&self, node_id_vec: &[TreeNodeID<Tree>]) -> &'t Tree::Node {
self.tree.get_node(self.get_lca_id(node_id_vec)).unwrap()
}
pub fn get_fa_index(&self, node_id: TreeNodeID<Tree>) -> usize {
let pos = self.fai[Into::<usize>::into(node_id)];
assert_ne!(
pos, NO_FIRST_APPEARANCE,
"node does not appear in the euler walk"
);
pos as usize
}
pub fn get_node_depth(&self, node_id: TreeNodeID<Tree>) -> usize {
self.da[self.get_fa_index(node_id)]
}
pub fn get_euler_pos(&self, pos: usize) -> TreeNodeID<Tree> {
self.euler[pos]
}
pub fn restrict_to_subtree<'c>(&self, child: &'c Tree) -> LcaOracle<'c, Tree> {
let subtree_root = child.get_root_id();
let start = self.get_fa_index(subtree_root);
let root_depth = self.da[start];
let mut end = start;
while end + 1 < self.euler.len() && self.da[end + 1] >= root_depth {
end += 1;
}
let euler: Vec<TreeNodeID<Tree>> = self.euler[start..=end].to_vec();
let da: Vec<usize> = self.da[start..=end]
.iter()
.map(|d| d - root_depth)
.collect();
let max_id: usize = euler
.iter()
.copied()
.map(Into::<usize>::into)
.max()
.expect("a subtree always contains at least its root");
let mut fai = vec![NO_FIRST_APPEARANCE; max_id + 1];
for (pos, node_id) in euler.iter().enumerate() {
let idx: usize = (*node_id).into();
if fai[idx] == NO_FIRST_APPEARANCE {
fai[idx] = pos as u32;
}
}
let rmq = FastRmq::from_vec(da.iter().map(|x| *x as u64).collect_vec());
LcaOracle {
tree: child,
euler,
fai,
da,
rmq,
}
}
pub fn euler_slice(&self) -> &[TreeNodeID<Tree>] {
&self.euler
}
pub fn depth_array(&self) -> &[usize] {
&self.da
}
pub fn heap_size(&self) -> usize {
self.euler.capacity() * std::mem::size_of::<TreeNodeID<Tree>>()
+ self.fai.capacity() * std::mem::size_of::<u32>()
+ self.da.capacity() * std::mem::size_of::<usize>()
+ self.rmq.heap_size()
}
}