use std::cell::{Cell, RefCell};
use std::collections::{HashMap, HashSet};
use std::rc::Rc;
use teksilo_core::signal::{ObserverHandle, Signal};
use crate::check_state::CheckState;
use crate::tree_change::NodeId;
use crate::tree_model::TreeModel;
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub enum AggregateMode {
None,
#[default]
DescendantsDriveAncestors,
}
struct Inner {
state: HashMap<NodeId, Signal<CheckState>>,
observers: HashMap<NodeId, ObserverHandle>,
bool_signals: HashMap<NodeId, Signal<bool>>,
bridge_guards: HashMap<NodeId, Rc<Cell<bool>>>,
bridge_observers: HashMap<NodeId, (ObserverHandle, ObserverHandle)>,
suppressed: HashSet<NodeId>,
}
pub struct TreeCheckedModel<T: 'static> {
tree: TreeModel<T>,
inner: Rc<RefCell<Inner>>,
mode: Rc<Cell<AggregateMode>>,
}
impl<T: 'static> TreeCheckedModel<T> {
pub fn new(tree: TreeModel<T>) -> Self {
Self {
tree,
inner: Rc::new(RefCell::new(Inner {
state: HashMap::new(),
observers: HashMap::new(),
bool_signals: HashMap::new(),
bridge_guards: HashMap::new(),
bridge_observers: HashMap::new(),
suppressed: HashSet::new(),
})),
mode: Rc::new(Cell::new(AggregateMode::default())),
}
}
pub fn with_mode(tree: TreeModel<T>, mode: AggregateMode) -> Self {
let m = Self::new(tree);
m.mode.set(mode);
m
}
pub fn aggregate_mode(&self) -> AggregateMode {
self.mode.get()
}
pub fn set_aggregate_mode(&self, mode: AggregateMode) {
self.mode.set(mode);
}
pub fn signal_for(&self, node: NodeId) -> Signal<CheckState> {
let sig = self
.inner
.borrow_mut()
.state
.entry(node)
.or_insert_with(|| Signal::new(CheckState::Unchecked))
.clone();
if !self.inner.borrow().observers.contains_key(&node) {
let handle = self.make_cascade_observer(&sig, node);
self.inner.borrow_mut().observers.insert(node, handle);
}
sig
}
fn make_cascade_observer(&self, sig: &Signal<CheckState>, node: NodeId) -> ObserverHandle {
let inner_w = Rc::downgrade(&self.inner);
let mode_w = Rc::downgrade(&self.mode);
let tree = self.tree.clone();
sig.observe(move |new_state| {
let inner_rc = match inner_w.upgrade() {
Some(rc) => rc,
None => return,
};
let mode_rc = match mode_w.upgrade() {
Some(rc) => rc,
None => return,
};
if inner_rc.borrow().suppressed.contains(&node) {
return;
}
if mode_rc.get() != AggregateMode::DescendantsDriveAncestors {
return;
}
if *new_state != CheckState::Indeterminate {
cascade_descendants(&tree, &inner_rc, node, *new_state);
}
let mut cur = tree.parent(node);
while let Some(p) = cur {
recompute_from_children(&tree, &inner_rc, p);
cur = tree.parent(p);
}
})
}
pub fn bool_signal_for(&self, node: NodeId) -> Signal<bool> {
if let Some(b) = self.inner.borrow().bool_signals.get(&node) {
return b.clone();
}
let tristate = self.signal_for(node);
let bool_sig = Signal::new(tristate.get() == CheckState::Checked);
let guard = Rc::new(Cell::new(false));
let bool_for_tri = bool_sig.clone();
let guard_for_tri = guard.clone();
let tri_to_bool = tristate.observe(move |state| {
if guard_for_tri.get() {
return;
}
let want = matches!(state, CheckState::Checked);
if bool_for_tri.get() != want {
guard_for_tri.set(true);
bool_for_tri.set(want);
guard_for_tri.set(false);
}
});
let tri_for_bool = tristate.clone();
let guard_for_bool = guard.clone();
let bool_to_tri = bool_sig.observe(move |checked| {
if guard_for_bool.get() {
return;
}
let want = if *checked {
CheckState::Checked
} else {
CheckState::Unchecked
};
if tri_for_bool.get() != want {
guard_for_bool.set(true);
tri_for_bool.set(want);
guard_for_bool.set(false);
}
});
let mut inner = self.inner.borrow_mut();
inner.bool_signals.insert(node, bool_sig.clone());
inner.bridge_guards.insert(node, guard);
inner
.bridge_observers
.insert(node, (tri_to_bool, bool_to_tri));
bool_sig
}
pub fn check_state(&self, node: NodeId) -> CheckState {
self.inner
.borrow()
.state
.get(&node)
.map(|s| s.get())
.unwrap_or(CheckState::Unchecked)
}
pub fn check(&self, node: NodeId) {
self.signal_for(node).set(CheckState::Checked);
}
pub fn uncheck(&self, node: NodeId) {
self.signal_for(node).set(CheckState::Unchecked);
}
pub fn toggle(&self, node: NodeId) {
let current = self.check_state(node);
let next = match (self.mode.get(), self.is_leaf(node), current) {
(AggregateMode::DescendantsDriveAncestors, true, CheckState::Unchecked) => {
CheckState::Checked
}
(AggregateMode::DescendantsDriveAncestors, true, _) => CheckState::Unchecked,
(_, _, _) => current.next_tristate(),
};
self.signal_for(node).set(next);
}
pub fn checked_nodes(&self) -> Vec<NodeId> {
self.inner
.borrow()
.state
.iter()
.filter_map(|(id, sig)| (sig.get() == CheckState::Checked).then_some(*id))
.collect()
}
pub fn clear(&self) {
let keys: Vec<NodeId> = self.inner.borrow().state.keys().copied().collect();
for k in keys {
write_state(&self.inner, k, CheckState::Unchecked);
}
}
fn is_leaf(&self, node: NodeId) -> bool {
self.tree.children(node).is_empty()
}
}
struct SuppressGuard {
inner: Rc<RefCell<Inner>>,
node: NodeId,
}
impl SuppressGuard {
fn new(inner: &Rc<RefCell<Inner>>, node: NodeId) -> Self {
inner.borrow_mut().suppressed.insert(node);
Self {
inner: inner.clone(),
node,
}
}
}
impl Drop for SuppressGuard {
fn drop(&mut self) {
if let Ok(mut inner) = self.inner.try_borrow_mut() {
inner.suppressed.remove(&self.node);
}
}
}
fn cascade_descendants<T: 'static>(
tree: &TreeModel<T>,
inner: &Rc<RefCell<Inner>>,
root: NodeId,
target: CheckState,
) {
for child in tree.children(root) {
write_state(inner, child, target);
cascade_descendants(tree, inner, child, target);
}
}
fn recompute_from_children<T: 'static>(
tree: &TreeModel<T>,
inner: &Rc<RefCell<Inner>>,
node: NodeId,
) {
let kids = tree.children(node);
if kids.is_empty() {
return;
}
let mut all_checked = true;
let mut all_unchecked = true;
for child in &kids {
let st = read_state(inner, *child);
match st {
CheckState::Checked => all_unchecked = false,
CheckState::Unchecked => all_checked = false,
CheckState::Indeterminate => {
all_checked = false;
all_unchecked = false;
}
}
}
let new_state = if all_checked {
CheckState::Checked
} else if all_unchecked {
CheckState::Unchecked
} else {
CheckState::Indeterminate
};
write_state(inner, node, new_state);
}
fn read_state(inner: &Rc<RefCell<Inner>>, node: NodeId) -> CheckState {
inner
.borrow()
.state
.get(&node)
.map(|s| s.get())
.unwrap_or(CheckState::Unchecked)
}
fn write_state(inner: &Rc<RefCell<Inner>>, node: NodeId, state: CheckState) {
let sig = {
let mut map = inner.borrow_mut();
map.state
.entry(node)
.or_insert_with(|| Signal::new(CheckState::Unchecked))
.clone()
};
if sig.get() != state {
let _guard = SuppressGuard::new(inner, node);
sig.set(state);
}
}
impl<T: 'static> Clone for TreeCheckedModel<T> {
fn clone(&self) -> Self {
Self {
tree: self.tree.clone(),
inner: self.inner.clone(),
mode: self.mode.clone(),
}
}
}
impl<T: 'static> std::fmt::Debug for TreeCheckedModel<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("TreeCheckedModel")
.field("mode", &self.mode.get())
.field("tracked_nodes", &self.inner.borrow().state.len())
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
fn sample_tree() -> (
TreeModel<&'static str>,
NodeId,
NodeId,
NodeId,
NodeId,
NodeId,
) {
let t = TreeModel::new();
let root1 = t.insert_root(0, "root1");
let a = t.insert_child(root1, 0, "a");
let b = t.insert_child(root1, 1, "b");
let root2 = t.insert_root(1, "root2");
let c = t.insert_child(root2, 0, "c");
(t, root1, a, b, root2, c)
}
#[test]
fn descendants_drive_ancestors_default() {
let (t, root1, a, b, _root2, _c) = sample_tree();
let m = TreeCheckedModel::new(t);
let _ = (m.signal_for(root1), m.signal_for(a), m.signal_for(b));
m.check(a);
assert_eq!(m.check_state(a), CheckState::Checked);
assert_eq!(m.check_state(root1), CheckState::Indeterminate);
m.check(b);
assert_eq!(m.check_state(root1), CheckState::Checked);
}
#[test]
fn set_parent_cascades_to_descendants() {
let (t, root1, a, b, _r2, _c) = sample_tree();
let m = TreeCheckedModel::new(t);
let _ = (m.signal_for(root1), m.signal_for(a), m.signal_for(b));
m.check(root1);
assert_eq!(m.check_state(a), CheckState::Checked);
assert_eq!(m.check_state(b), CheckState::Checked);
assert_eq!(m.check_state(root1), CheckState::Checked);
m.uncheck(root1);
assert_eq!(m.check_state(a), CheckState::Unchecked);
assert_eq!(m.check_state(b), CheckState::Unchecked);
}
#[test]
fn external_signal_write_triggers_cascade() {
let (t, root1, a, b, _r2, _c) = sample_tree();
let m = TreeCheckedModel::new(t);
let parent_sig = m.signal_for(root1);
let _ = (m.signal_for(a), m.signal_for(b));
parent_sig.set(CheckState::Checked);
assert_eq!(m.check_state(a), CheckState::Checked);
assert_eq!(m.check_state(b), CheckState::Checked);
}
#[test]
fn lazy_signal_still_cascades() {
let (t, root1, a, _b, _r2, _c) = sample_tree();
let m = TreeCheckedModel::new(t);
let _ = m.signal_for(root1); m.check(root1); assert_eq!(m.check_state(a), CheckState::Checked);
let a_sig = m.signal_for(a); a_sig.set(CheckState::Unchecked); assert_eq!(m.check_state(root1), CheckState::Indeterminate);
}
#[test]
fn aggregate_mode_none_disables_propagation() {
let (t, root1, a, _b, _r2, _c) = sample_tree();
let m = TreeCheckedModel::with_mode(t, AggregateMode::None);
let _ = (m.signal_for(root1), m.signal_for(a));
m.check(a);
assert_eq!(m.check_state(a), CheckState::Checked);
assert_eq!(m.check_state(root1), CheckState::Unchecked);
}
#[test]
fn signal_for_is_stable_across_calls() {
let (t, _root1, a, _b, _r2, _c) = sample_tree();
let m = TreeCheckedModel::new(t);
let s1 = m.signal_for(a);
let s2 = m.signal_for(a);
m.check(a);
assert_eq!(s1.get(), CheckState::Checked);
assert_eq!(s2.get(), CheckState::Checked);
}
#[test]
fn checked_nodes_excludes_indeterminate() {
let (t, root1, a, _b, _r2, _c) = sample_tree();
let m = TreeCheckedModel::new(t);
let _ = (m.signal_for(root1), m.signal_for(a));
m.check(a);
let nodes = m.checked_nodes();
assert!(nodes.contains(&a));
assert!(!nodes.contains(&root1));
}
#[test]
fn toggle_leaf_two_state_in_aggregate_mode() {
let (t, _root1, a, _b, _r2, _c) = sample_tree();
let m = TreeCheckedModel::new(t);
m.toggle(a);
assert_eq!(m.check_state(a), CheckState::Checked);
m.toggle(a);
assert_eq!(m.check_state(a), CheckState::Unchecked);
}
#[test]
fn bool_signal_writes_propagate_to_tristate() {
let (t, _root1, a, _b, _r2, _c) = sample_tree();
let m = TreeCheckedModel::new(t);
let bool_sig = m.bool_signal_for(a);
assert!(!bool_sig.get());
bool_sig.set(true);
assert_eq!(m.check_state(a), CheckState::Checked);
bool_sig.set(false);
assert_eq!(m.check_state(a), CheckState::Unchecked);
}
#[test]
fn bool_signal_reflects_tristate_writes() {
let (t, _root1, a, _b, _r2, _c) = sample_tree();
let m = TreeCheckedModel::new(t);
let bool_sig = m.bool_signal_for(a);
m.check(a);
assert!(bool_sig.get());
m.uncheck(a);
assert!(!bool_sig.get());
}
#[test]
fn bool_signal_indeterminate_reads_as_false() {
let (t, root1, a, b, _r2, _c) = sample_tree();
let m = TreeCheckedModel::new(t);
let parent_bool = m.bool_signal_for(root1);
let _ = (m.signal_for(a), m.signal_for(b));
m.check(a); assert_eq!(m.check_state(root1), CheckState::Indeterminate);
assert!(!parent_bool.get(), "Indeterminate must not read as true");
}
#[test]
fn bool_signal_writes_through_leaves_recompute_ancestors() {
let (t, root1, a, b, _r2, _c) = sample_tree();
let m = TreeCheckedModel::new(t);
let a_bool = m.bool_signal_for(a);
let b_bool = m.bool_signal_for(b);
a_bool.set(true);
assert_eq!(m.check_state(root1), CheckState::Indeterminate);
b_bool.set(true);
assert_eq!(m.check_state(root1), CheckState::Checked);
}
#[test]
fn bool_signal_for_is_stable_across_calls() {
let (t, _root1, a, _b, _r2, _c) = sample_tree();
let m = TreeCheckedModel::new(t);
let s1 = m.bool_signal_for(a);
let s2 = m.bool_signal_for(a);
s1.set(true);
assert!(s2.get());
}
#[test]
fn clear_resets_all() {
let (t, root1, _a, _b, _r2, _c) = sample_tree();
let m = TreeCheckedModel::new(t);
let _ = (m.signal_for(root1),);
m.check(root1);
m.clear();
assert_eq!(m.checked_nodes(), Vec::<NodeId>::new());
}
#[test]
fn clear_resets_every_node_and_still_notifies() {
let (t, root1, a, b, root2, c) = sample_tree();
let m = TreeCheckedModel::new(t);
let _ = (m.signal_for(root1), m.signal_for(a), m.signal_for(b));
let bool_a = m.bool_signal_for(a);
m.check(root1); m.check(c); assert_eq!(m.check_state(root1), CheckState::Checked);
assert_eq!(m.check_state(root2), CheckState::Checked);
assert!(bool_a.get());
let notified: Rc<RefCell<HashSet<NodeId>>> = Rc::new(RefCell::new(HashSet::new()));
let mut handles = Vec::new();
for node in [root1, a, b, root2, c] {
let log = notified.clone();
handles.push(m.signal_for(node).observe(move |state| {
if *state == CheckState::Unchecked {
log.borrow_mut().insert(node);
}
}));
}
m.clear();
assert_eq!(m.checked_nodes(), Vec::<NodeId>::new());
for node in [root1, a, b, root2, c] {
assert_eq!(m.check_state(node), CheckState::Unchecked);
}
assert!(!bool_a.get());
assert_eq!(
notified.borrow().len(),
5,
"every tracked node must still notify its own observers on clear: {:?}",
notified.borrow()
);
drop(handles);
}
#[test]
fn reentrant_write_to_unrelated_node_still_cascades() {
let (t, root1, a, b, root2, c) = sample_tree();
let m = TreeCheckedModel::new(t);
let _ = (
m.signal_for(root1),
m.signal_for(a),
m.signal_for(b),
m.signal_for(root2),
m.signal_for(c),
);
let m_for_observer = m.clone();
let _obs = m.signal_for(a).observe(move |state| {
if *state == CheckState::Checked {
m_for_observer.check(c);
}
});
m.check(root1);
assert_eq!(m.check_state(a), CheckState::Checked);
assert_eq!(m.check_state(c), CheckState::Checked);
assert_eq!(m.check_state(root2), CheckState::Checked);
}
}