use std::cmp::Ordering;
use std::marker::PhantomData;
use crate::dnd_types::ItemKey;
use crate::sort_filter_tree_model::TreeFilterMode;
use crate::tree_data_slice::TreeRow;
type Predicate<T> = Box<dyn Fn(&T) -> bool>;
type Comparator<T> = Box<dyn Fn(&T, &T) -> Ordering>;
pub struct TreeRowFilter<K: ItemKey, T> {
predicate: Option<Predicate<T>>,
mode: TreeFilterMode,
comparator: Option<Comparator<T>>,
_k: PhantomData<fn() -> K>,
}
impl<K: ItemKey, T: 'static> Default for TreeRowFilter<K, T> {
fn default() -> Self {
Self::new()
}
}
impl<K: ItemKey, T: 'static> TreeRowFilter<K, T> {
pub fn new() -> Self {
Self {
predicate: None,
mode: TreeFilterMode::default(),
comparator: None,
_k: PhantomData,
}
}
pub fn filter_mode(mut self, mode: TreeFilterMode) -> Self {
self.mode = mode;
self
}
pub fn filter(mut self, pred: impl Fn(&T) -> bool + 'static) -> Self {
self.predicate = Some(Box::new(pred));
self
}
pub fn sort(mut self, cmp: impl Fn(&T, &T) -> Ordering + 'static) -> Self {
self.comparator = Some(Box::new(cmp));
self
}
pub fn sort_desc(mut self, cmp: impl Fn(&T, &T) -> Ordering + 'static) -> Self {
self.comparator = Some(Box::new(move |a, b| cmp(a, b).reverse()));
self
}
pub fn apply(&self, rows: Vec<TreeRow<K, T>>) -> Vec<TreeRow<K, T>> {
if self.predicate.is_none() && self.comparator.is_none() {
return rows;
}
let n = rows.len();
let mut children: Vec<Vec<usize>> = vec![Vec::new(); n];
let mut parent_of: Vec<Option<usize>> = vec![None; n];
let mut roots: Vec<usize> = Vec::new();
let mut stack: Vec<(usize, usize)> = Vec::new(); for (i, row) in rows.iter().enumerate() {
while let Some(&(d, _)) = stack.last() {
if d >= row.depth {
stack.pop();
} else {
break;
}
}
match stack.last() {
Some(&(_, parent)) => {
children[parent].push(i);
parent_of[i] = Some(parent);
}
None => roots.push(i),
}
stack.push((row.depth, i));
}
let visible = self.compute_visible(&rows, &children, &roots, &parent_of);
if let Some(cmp) = &self.comparator {
roots.sort_by(|&a, &b| cmp(&rows[a].item, &rows[b].item));
for list in children.iter_mut() {
list.sort_by(|&a, &b| cmp(&rows[a].item, &rows[b].item));
}
}
let mut emit: Vec<(usize, usize)> = Vec::with_capacity(n);
for &root in &roots {
emit_dfs(root, 0, &children, &visible, &mut emit);
}
let mut slots: Vec<Option<TreeRow<K, T>>> = rows.into_iter().map(Some).collect();
emit.into_iter()
.map(|(i, depth)| {
let mut row = slots[i].take().expect("each node emitted at most once");
row.depth = depth;
row
})
.collect()
}
fn compute_visible(
&self,
rows: &[TreeRow<K, T>],
children: &[Vec<usize>],
roots: &[usize],
parent_of: &[Option<usize>],
) -> Vec<bool> {
let Some(pred) = &self.predicate else {
return vec![true; rows.len()];
};
let matches: Vec<bool> = rows.iter().map(|r| pred(&r.item)).collect();
let mut visible = vec![false; rows.len()];
match self.mode {
TreeFilterMode::HideNonMatching => {
for i in 0..rows.len() {
visible[i] = matches[i] && parent_of[i].is_none_or(|p| visible[p]);
}
}
TreeFilterMode::KeepAncestors => {
for &r in roots {
keep_ancestors(r, children, &matches, &mut visible);
}
}
TreeFilterMode::KeepDescendants => {
for &r in roots {
keep_descendants(r, children, &matches, &mut visible);
}
}
}
visible
}
}
fn keep_ancestors(root: usize, children: &[Vec<usize>], matches: &[bool], visible: &mut [bool]) {
let mut pre_order = Vec::with_capacity(children.len());
let mut stack = vec![root];
while let Some(i) = stack.pop() {
pre_order.push(i);
for &c in children[i].iter().rev() {
stack.push(c);
}
}
for &i in pre_order.iter().rev() {
let any_descendant = children[i].iter().any(|&c| visible[c]);
if matches[i] || any_descendant {
visible[i] = true;
}
}
}
fn keep_descendants(root: usize, children: &[Vec<usize>], matches: &[bool], visible: &mut [bool]) {
let mut stack = vec![(root, false)];
while let Some((i, ancestor_matched)) = stack.pop() {
let here = matches[i] || ancestor_matched;
if here {
visible[i] = true;
}
for &c in children[i].iter().rev() {
stack.push((c, here));
}
}
}
fn emit_dfs(
root: usize,
out_depth: usize,
children: &[Vec<usize>],
visible: &[bool],
emit: &mut Vec<(usize, usize)>,
) {
let mut stack = vec![(root, out_depth)];
while let Some((i, out_depth)) = stack.pop() {
let child_depth = if visible[i] {
emit.push((i, out_depth));
out_depth + 1
} else {
out_depth
};
for &c in children[i].iter().rev() {
stack.push((c, child_depth));
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn sample() -> Vec<TreeRow<u64, &'static str>> {
vec![
TreeRow::new(1, "Manuscript", 0),
TreeRow::new(2, "Book One", 1),
TreeRow::new(3, "Opening", 2),
TreeRow::new(4, "Dawn", 2),
TreeRow::new(5, "Chapter Two", 1),
TreeRow::new(6, "Fight", 2),
TreeRow::new(7, "Notes", 0),
TreeRow::new(8, "Sketch", 1),
]
}
fn titles(rows: &[TreeRow<u64, &'static str>]) -> Vec<&'static str> {
rows.iter().map(|r| r.item).collect()
}
#[test]
fn identity_passes_through() {
let out = TreeRowFilter::new().apply(sample());
assert_eq!(out.len(), 8);
assert_eq!(titles(&out), titles(&sample()));
}
#[test]
fn keep_ancestors_shows_path_to_match() {
let out = TreeRowFilter::new()
.filter_mode(TreeFilterMode::KeepAncestors)
.filter(|t: &&str| *t == "Dawn")
.apply(sample());
assert_eq!(titles(&out), vec!["Manuscript", "Book One", "Dawn"]);
assert_eq!(
out.iter().map(|r| r.depth).collect::<Vec<_>>(),
vec![0, 1, 2]
);
}
#[test]
fn keep_descendants_surfaces_subtree_even_under_nonmatching_ancestor() {
let out = TreeRowFilter::new()
.filter_mode(TreeFilterMode::KeepDescendants)
.filter(|t: &&str| *t == "Book One")
.apply(sample());
assert_eq!(titles(&out), vec!["Book One", "Opening", "Dawn"]);
assert_eq!(
out.iter().map(|r| r.depth).collect::<Vec<_>>(),
vec![0, 1, 1]
);
}
#[test]
fn hide_non_matching_requires_whole_path() {
let out = TreeRowFilter::new()
.filter_mode(TreeFilterMode::HideNonMatching)
.filter(|t: &&str| *t == "Manuscript" || *t == "Book One")
.apply(sample());
assert_eq!(titles(&out), vec!["Manuscript", "Book One"]);
assert_eq!(out.iter().map(|r| r.depth).collect::<Vec<_>>(), vec![0, 1]);
}
#[test]
fn hide_non_matching_hides_match_under_hidden_parent() {
let out = TreeRowFilter::new()
.filter_mode(TreeFilterMode::HideNonMatching)
.filter(|t: &&str| *t == "Opening")
.apply(sample());
assert!(out.is_empty());
}
#[test]
fn empty_match_yields_empty() {
let out = TreeRowFilter::new()
.filter_mode(TreeFilterMode::KeepAncestors)
.filter(|_: &&str| false)
.apply(sample());
assert!(out.is_empty());
}
#[test]
fn sort_reorders_siblings_per_parent() {
let out = TreeRowFilter::new()
.sort(|a: &&str, b: &&str| a.cmp(b))
.apply(sample());
assert_eq!(
titles(&out),
vec![
"Manuscript",
"Book One",
"Dawn",
"Opening",
"Chapter Two",
"Fight",
"Notes",
"Sketch"
]
);
}
#[test]
fn sort_desc_reverses() {
let out = TreeRowFilter::new()
.sort_desc(|a: &&str, b: &&str| a.cmp(b))
.apply(sample());
assert_eq!(out[0].item, "Notes");
assert_eq!(out[1].item, "Sketch");
assert_eq!(out[2].item, "Manuscript");
}
#[test]
fn filter_then_sort_compose() {
let out = TreeRowFilter::new()
.filter_mode(TreeFilterMode::KeepAncestors)
.filter(|t: &&str| *t == "Dawn" || *t == "Opening")
.sort(|a: &&str, b: &&str| a.cmp(b))
.apply(sample());
assert_eq!(
titles(&out),
vec!["Manuscript", "Book One", "Dawn", "Opening"]
);
}
#[test]
fn structure_preserved_when_all_match() {
let out = TreeRowFilter::new()
.filter_mode(TreeFilterMode::KeepAncestors)
.filter(|_: &&str| true)
.apply(sample());
assert_eq!(titles(&out), titles(&sample()));
assert_eq!(
out.iter().map(|r| r.depth).collect::<Vec<_>>(),
vec![0, 1, 2, 2, 1, 2, 0, 1]
);
}
#[test]
fn deep_chain_applies_each_mode_without_overflow() {
const DEPTH: usize = 50_000;
let rows = |depth: usize| -> Vec<TreeRow<u64, usize>> {
(0..depth).map(|i| TreeRow::new(i as u64, i, i)).collect()
};
let out = TreeRowFilter::new()
.filter_mode(TreeFilterMode::KeepAncestors)
.filter(move |item: &usize| *item == DEPTH - 1)
.apply(rows(DEPTH));
assert_eq!(out.len(), DEPTH);
assert_eq!(out.last().unwrap().depth, DEPTH - 1);
let out = TreeRowFilter::new()
.filter_mode(TreeFilterMode::KeepDescendants)
.filter(|item: &usize| *item == 0)
.apply(rows(DEPTH));
assert_eq!(out.len(), DEPTH);
let out = TreeRowFilter::new()
.filter_mode(TreeFilterMode::HideNonMatching)
.filter(|item: &usize| *item == 0)
.apply(rows(DEPTH));
assert_eq!(out.len(), 1);
}
}