use serde::{Deserialize, Serialize};
use std::collections::HashMap;
pub trait TreeNode {
fn id(&self) -> i64;
fn parent_id(&self) -> i64;
fn children(&self) -> &Vec<Self>
where
Self: Sized;
fn children_mut(&mut self) -> &mut Vec<Self>
where
Self: Sized;
}
pub struct TreeUtils;
impl TreeUtils {
pub fn build_tree<T>(nodes: Vec<T>) -> Vec<T>
where
T: TreeNode + Clone,
{
let mut node_map: HashMap<i64, T> = HashMap::new();
let mut root_nodes: Vec<T> = Vec::new();
for node in nodes.iter() {
node_map.insert(node.id(), node.clone());
}
for node in nodes {
let parent_id = node.parent_id();
if parent_id == 0 {
root_nodes.push(node);
} else if let Some(parent) = node_map.get_mut(&parent_id) {
parent.children_mut().push(node);
}
}
root_nodes
}
pub fn get_child_ids<T>(nodes: &[T], parent_id: i64) -> Vec<i64>
where
T: TreeNode,
{
let mut child_ids = Vec::new();
for node in nodes {
if node.parent_id() == parent_id {
child_ids.push(node.id());
child_ids.extend(Self::get_child_ids(nodes, node.id()));
}
}
child_ids
}
pub fn get_parent_ids<T>(nodes: &[T], child_id: i64) -> Vec<i64>
where
T: TreeNode,
{
let mut parent_ids = Vec::new();
if let Some(node) = nodes.iter().find(|n| n.id() == child_id) {
let parent_id = node.parent_id();
if parent_id != 0 {
parent_ids.push(parent_id);
parent_ids.extend(Self::get_parent_ids(nodes, parent_id));
}
}
parent_ids
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CommonTreeNode<T> {
pub id: i64,
pub parent_id: i64,
pub data: T,
pub children: Vec<CommonTreeNode<T>>,
}
impl<T> TreeNode for CommonTreeNode<T> {
fn id(&self) -> i64 {
self.id
}
fn parent_id(&self) -> i64 {
self.parent_id
}
fn children(&self) -> &Vec<Self> {
&self.children
}
fn children_mut(&mut self) -> &mut Vec<Self> {
&mut self.children
}
}