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::dnd_types::ItemKey;
use crate::tree_checked_model::AggregateMode;
use crate::tree_data_source::TreeDataSource;
type ChildrenFn<K> = Rc<dyn Fn(&K) -> Vec<K>>;
type ParentFn<K> = Rc<dyn Fn(&K) -> Option<K>>;
struct Inner<K: ItemKey> {
state: HashMap<K, Signal<CheckState>>,
observers: HashMap<K, ObserverHandle>,
bool_signals: HashMap<K, Signal<bool>>,
bridge_guards: HashMap<K, Rc<Cell<bool>>>,
bridge_observers: HashMap<K, (ObserverHandle, ObserverHandle)>,
suppressed: HashSet<K>,
}
pub struct KeyedTreeCheckedModel<K: ItemKey> {
children: ChildrenFn<K>,
parent: ParentFn<K>,
inner: Rc<RefCell<Inner<K>>>,
mode: Rc<Cell<AggregateMode>>,
}
impl<K: ItemKey> KeyedTreeCheckedModel<K> {
pub fn new(
children: impl Fn(&K) -> Vec<K> + 'static,
parent: impl Fn(&K) -> Option<K> + 'static,
) -> Self {
Self {
children: Rc::new(children),
parent: Rc::new(parent),
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 from_source<S>(source: S) -> Self
where
S: TreeDataSource<Key = K> + Clone + 'static,
{
let for_children = source.clone();
let for_parent = source;
Self::new(
move |k| for_children.child_keys(k),
move |k| for_parent.parent(k),
)
}
pub fn with_mode(self, mode: AggregateMode) -> Self {
self.mode.set(mode);
self
}
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, key: K) -> Signal<CheckState> {
let sig = self
.inner
.borrow_mut()
.state
.entry(key.clone())
.or_insert_with(|| Signal::new(CheckState::Unchecked))
.clone();
if !self.inner.borrow().observers.contains_key(&key) {
let handle = self.make_cascade_observer(&sig, key.clone());
self.inner.borrow_mut().observers.insert(key, handle);
}
sig
}
fn make_cascade_observer(&self, sig: &Signal<CheckState>, node: K) -> ObserverHandle {
let inner_w = Rc::downgrade(&self.inner);
let mode_w = Rc::downgrade(&self.mode);
let children = self.children.clone();
let parent = self.parent.clone();
sig.observe(move |new_state| {
let Some(inner_rc) = inner_w.upgrade() else {
return;
};
let Some(mode_rc) = mode_w.upgrade() else {
return;
};
if inner_rc.borrow().suppressed.contains(&node) {
return;
}
if mode_rc.get() != AggregateMode::DescendantsDriveAncestors {
return;
}
if *new_state != CheckState::Indeterminate {
cascade_descendants(&children, &inner_rc, &node, *new_state);
}
let mut cur = parent(&node);
while let Some(p) = cur {
recompute_from_children(&children, &inner_rc, &p);
cur = parent(&p);
}
})
}
pub fn bool_signal_for(&self, key: K) -> Signal<bool> {
if let Some(b) = self.inner.borrow().bool_signals.get(&key) {
return b.clone();
}
let tristate = self.signal_for(key.clone());
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(key.clone(), bool_sig.clone());
inner.bridge_guards.insert(key.clone(), guard);
inner
.bridge_observers
.insert(key, (tri_to_bool, bool_to_tri));
bool_sig
}
pub fn check_state(&self, key: &K) -> CheckState {
self.inner
.borrow()
.state
.get(key)
.map(|s| s.get())
.unwrap_or(CheckState::Unchecked)
}
pub fn check(&self, key: K) {
self.signal_for(key).set(CheckState::Checked);
}
pub fn uncheck(&self, key: K) {
self.signal_for(key).set(CheckState::Unchecked);
}
pub fn toggle(&self, key: K) {
let current = self.check_state(&key);
let next = match (self.mode.get(), self.is_leaf(&key), current) {
(AggregateMode::DescendantsDriveAncestors, true, CheckState::Unchecked) => {
CheckState::Checked
}
(AggregateMode::DescendantsDriveAncestors, true, _) => CheckState::Unchecked,
(_, _, _) => current.next_tristate(),
};
self.signal_for(key).set(next);
}
pub fn checked_keys(&self) -> Vec<K> {
self.inner
.borrow()
.state
.iter()
.filter(|(_, sig)| sig.get() == CheckState::Checked)
.map(|(k, _)| k.clone())
.collect()
}
pub fn clear(&self) {
let keys: Vec<K> = self.inner.borrow().state.keys().cloned().collect();
for k in keys {
write_state(&self.inner, &k, CheckState::Unchecked);
}
}
pub fn prune_missing(&self, exists: impl Fn(&K) -> bool) {
let all: Vec<K> = self.inner.borrow().state.keys().cloned().collect();
let stale: Vec<K> = all.into_iter().filter(|k| !exists(k)).collect();
if !stale.is_empty() {
let mut inner = self.inner.borrow_mut();
for k in &stale {
inner.state.remove(k);
inner.observers.remove(k);
inner.bool_signals.remove(k);
inner.bridge_guards.remove(k);
inner.bridge_observers.remove(k);
}
}
self.reaggregate();
}
pub fn reaggregate(&self) {
if self.mode.get() != AggregateMode::DescendantsDriveAncestors {
return;
}
let tracked: Vec<K> = self.inner.borrow().state.keys().cloned().collect();
let mut keys: HashSet<K> = tracked.iter().cloned().collect();
for k in &tracked {
let mut cur = (self.parent)(k);
let mut guard = 0usize;
while let Some(p) = cur {
if !keys.insert(p.clone()) {
break;
}
guard += 1;
if guard > 1_000_000 {
break; }
cur = (self.parent)(&p);
}
}
let mut keys: Vec<K> = keys.into_iter().collect();
keys.sort_by_key(|k| std::cmp::Reverse(self.depth_of(k)));
for k in keys {
recompute_from_children(&self.children, &self.inner, &k);
}
}
fn depth_of(&self, key: &K) -> usize {
let mut depth = 0usize;
let mut cur = (self.parent)(key);
while let Some(p) = cur {
depth += 1;
if depth > 1_000_000 {
break;
}
cur = (self.parent)(&p);
}
depth
}
fn is_leaf(&self, key: &K) -> bool {
(self.children)(key).is_empty()
}
}
struct SuppressGuard<K: ItemKey> {
inner: Rc<RefCell<Inner<K>>>,
key: K,
}
impl<K: ItemKey> SuppressGuard<K> {
fn new(inner: &Rc<RefCell<Inner<K>>>, key: K) -> Self {
inner.borrow_mut().suppressed.insert(key.clone());
Self {
inner: inner.clone(),
key,
}
}
}
impl<K: ItemKey> Drop for SuppressGuard<K> {
fn drop(&mut self) {
if let Ok(mut inner) = self.inner.try_borrow_mut() {
inner.suppressed.remove(&self.key);
}
}
}
fn cascade_descendants<K: ItemKey>(
children: &ChildrenFn<K>,
inner: &Rc<RefCell<Inner<K>>>,
root: &K,
target: CheckState,
) {
for child in children(root) {
write_state(inner, &child, target);
cascade_descendants(children, inner, &child, target);
}
}
fn recompute_from_children<K: ItemKey>(
children: &ChildrenFn<K>,
inner: &Rc<RefCell<Inner<K>>>,
node: &K,
) {
let kids = children(node);
if kids.is_empty() {
return;
}
let mut all_checked = true;
let mut all_unchecked = true;
for child in &kids {
match read_state(inner, child) {
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<K: ItemKey>(inner: &Rc<RefCell<Inner<K>>>, node: &K) -> CheckState {
inner
.borrow()
.state
.get(node)
.map(|s| s.get())
.unwrap_or(CheckState::Unchecked)
}
fn write_state<K: ItemKey>(inner: &Rc<RefCell<Inner<K>>>, node: &K, state: CheckState) {
let sig = {
let mut map = inner.borrow_mut();
map.state
.entry(node.clone())
.or_insert_with(|| Signal::new(CheckState::Unchecked))
.clone()
};
if sig.get() != state {
let _guard = SuppressGuard::new(inner, node.clone());
sig.set(state);
}
}
impl<K: ItemKey> Clone for KeyedTreeCheckedModel<K> {
fn clone(&self) -> Self {
Self {
children: self.children.clone(),
parent: self.parent.clone(),
inner: self.inner.clone(),
mode: self.mode.clone(),
}
}
}
impl<K: ItemKey> std::fmt::Debug for KeyedTreeCheckedModel<K> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("KeyedTreeCheckedModel")
.field("mode", &self.mode.get())
.field("tracked_nodes", &self.inner.borrow().state.len())
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{TreeDataSlice, TreeRow};
fn slice() -> TreeDataSlice<u64, &'static str> {
TreeDataSlice::from_rows(vec![
TreeRow::new(1, "Binder", 0),
TreeRow::new(2, "Chapter", 1),
TreeRow::new(3, "Scene A", 2),
TreeRow::new(5, "Scene C", 2),
TreeRow::new(4, "Scene B", 1),
])
}
fn model() -> KeyedTreeCheckedModel<u64> {
let m = KeyedTreeCheckedModel::from_source(slice());
let _ = (
m.signal_for(1),
m.signal_for(2),
m.signal_for(3),
m.signal_for(4),
m.signal_for(5),
);
m
}
#[test]
fn descendants_drive_ancestors() {
let m = model();
m.check(3);
assert_eq!(m.check_state(&3), CheckState::Checked);
assert_eq!(m.check_state(&2), CheckState::Indeterminate);
assert_eq!(m.check_state(&1), CheckState::Indeterminate);
m.check(5);
assert_eq!(m.check_state(&2), CheckState::Checked);
assert_eq!(m.check_state(&1), CheckState::Indeterminate); m.check(4);
assert_eq!(m.check_state(&1), CheckState::Checked);
}
#[test]
fn parent_cascades_to_descendants() {
let m = model();
m.check(2); assert_eq!(m.check_state(&3), CheckState::Checked);
assert_eq!(m.check_state(&5), CheckState::Checked);
m.uncheck(2);
assert_eq!(m.check_state(&3), CheckState::Unchecked);
assert_eq!(m.check_state(&5), CheckState::Unchecked);
}
#[test]
fn check_root_cascades_whole_tree() {
let m = model();
m.check(1);
for k in [2u64, 3, 4, 5] {
assert_eq!(m.check_state(&k), CheckState::Checked);
}
}
#[test]
fn aggregate_mode_none_is_independent() {
let m = KeyedTreeCheckedModel::from_source(slice()).with_mode(AggregateMode::None);
let _ = (m.signal_for(1), m.signal_for(3));
m.check(3);
assert_eq!(m.check_state(&3), CheckState::Checked);
assert_eq!(m.check_state(&2), CheckState::Unchecked);
assert_eq!(m.check_state(&1), CheckState::Unchecked);
}
#[test]
fn external_signal_write_cascades() {
let m = model();
m.signal_for(2).set(CheckState::Checked); assert_eq!(m.check_state(&3), CheckState::Checked);
assert_eq!(m.check_state(&5), CheckState::Checked);
}
#[test]
fn bool_signal_bridge() {
let m = model();
let b = m.bool_signal_for(3);
assert!(!b.get());
b.set(true);
assert_eq!(m.check_state(&3), CheckState::Checked);
m.uncheck(3);
assert!(!b.get());
}
#[test]
fn bool_signal_indeterminate_reads_false() {
let m = model();
let binder_bool = m.bool_signal_for(1);
m.check(3); assert_eq!(m.check_state(&1), CheckState::Indeterminate);
assert!(!binder_bool.get());
}
#[test]
fn checked_keys_excludes_indeterminate() {
let m = model();
m.check(3);
let keys = m.checked_keys();
assert!(keys.contains(&3));
assert!(!keys.contains(&2)); assert!(!keys.contains(&1));
}
#[test]
fn signal_is_stable_across_calls() {
let m = model();
let s1 = m.signal_for(3);
let s2 = m.signal_for(3);
m.check(3);
assert_eq!(s1.get(), CheckState::Checked);
assert_eq!(s2.get(), CheckState::Checked);
}
#[test]
fn prune_missing_drops_stale_state() {
let m = model();
m.check(3);
assert!(m.checked_keys().contains(&3));
m.prune_missing(|k| *k != 3);
assert!(!m.checked_keys().contains(&3));
assert_eq!(m.check_state(&3), CheckState::Unchecked); }
#[test]
fn clear_resets_all() {
let m = model();
m.check(1);
m.clear();
assert_eq!(m.checked_keys(), Vec::<u64>::new());
}
#[test]
fn clear_resets_every_key_and_still_notifies() {
let m = model();
m.check(2);
m.check(4);
assert_eq!(m.check_state(&1), CheckState::Checked);
assert_eq!(m.check_state(&2), CheckState::Checked);
assert_eq!(m.check_state(&3), CheckState::Checked);
assert_eq!(m.check_state(&5), CheckState::Checked);
let bool_3 = m.bool_signal_for(3);
assert!(bool_3.get());
let notified: Rc<RefCell<HashSet<u64>>> = Rc::new(RefCell::new(HashSet::new()));
let mut handles = Vec::new();
for key in [1u64, 2, 3, 4, 5] {
let log = notified.clone();
handles.push(m.signal_for(key).observe(move |state| {
if *state == CheckState::Unchecked {
log.borrow_mut().insert(key);
}
}));
}
m.clear();
assert_eq!(m.checked_keys(), Vec::<u64>::new());
for key in [1u64, 2, 3, 4, 5] {
assert_eq!(m.check_state(&key), CheckState::Unchecked);
}
assert!(!bool_3.get());
assert_eq!(
notified.borrow().len(),
5,
"every tracked key must still notify its own observers on clear: {:?}",
notified.borrow()
);
drop(handles);
}
#[test]
fn lazy_signal_still_cascades() {
let m = KeyedTreeCheckedModel::from_source(slice());
let _ = m.signal_for(1); m.check(1); assert_eq!(m.check_state(&3), CheckState::Checked);
let scene3 = m.signal_for(3); scene3.set(CheckState::Unchecked); assert_eq!(m.check_state(&2), CheckState::Indeterminate);
assert_eq!(m.check_state(&1), CheckState::Indeterminate);
}
#[test]
fn prune_missing_reaggregates_ancestors() {
let s = slice(); let m = KeyedTreeCheckedModel::from_source(s.clone());
let _ = (
m.signal_for(1),
m.signal_for(2),
m.signal_for(3),
m.signal_for(5),
);
m.check(3);
assert_eq!(m.check_state(&2), CheckState::Indeterminate);
assert_eq!(m.check_state(&1), CheckState::Indeterminate);
s.set_rows(vec![
TreeRow::new(1, "Binder", 0),
TreeRow::new(2, "Chapter", 1),
TreeRow::new(5, "Scene C", 2),
TreeRow::new(4, "Scene B", 1),
]);
assert_eq!(s.child_keys_of(&2), vec![5]);
m.prune_missing(|k| s.contains_key(k));
assert!(!m.checked_keys().contains(&3));
assert_eq!(m.check_state(&2), CheckState::Unchecked);
assert_eq!(m.check_state(&1), CheckState::Unchecked);
}
#[test]
fn reaggregate_after_added_child() {
let s = slice();
let m = KeyedTreeCheckedModel::from_source(s.clone());
let _ = (m.signal_for(2), m.signal_for(3), m.signal_for(5));
m.check(2); assert_eq!(m.check_state(&2), CheckState::Checked);
s.set_rows(vec![
TreeRow::new(1, "Binder", 0),
TreeRow::new(2, "Chapter", 1),
TreeRow::new(3, "Scene A", 2),
TreeRow::new(5, "Scene C", 2),
TreeRow::new(6, "Scene D", 2),
TreeRow::new(4, "Scene B", 1),
]);
m.reaggregate();
assert_eq!(m.check_state(&2), CheckState::Indeterminate);
}
#[test]
fn new_with_explicit_closures() {
let parents: HashMap<&str, &str> = [("B", "A"), ("C", "A")].into_iter().collect();
let m = KeyedTreeCheckedModel::<&str>::new(
|k| match *k {
"A" => vec!["B", "C"],
_ => vec![],
},
move |k| parents.get(k).copied(),
);
let _ = (m.signal_for("A"), m.signal_for("B"), m.signal_for("C"));
m.check("B");
assert_eq!(m.check_state(&"A"), CheckState::Indeterminate);
m.check("C");
assert_eq!(m.check_state(&"A"), CheckState::Checked);
}
#[test]
fn reentrant_write_to_unrelated_key_still_cascades() {
let parents: HashMap<&str, &str> = [("a", "root1"), ("b", "root1"), ("c", "root2")]
.into_iter()
.collect();
let m = KeyedTreeCheckedModel::<&str>::new(
|k| match *k {
"root1" => vec!["a", "b"],
"root2" => vec!["c"],
_ => vec![],
},
move |k| parents.get(k).copied(),
);
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);
}
}