Skip to main content

cranpose_ui/modifier/
local.rs

1use std::any::{Any, TypeId};
2use std::collections::{HashMap, HashSet};
3use std::fmt;
4use std::hash::{Hash, Hasher};
5use std::rc::Rc;
6
7use cranpose_foundation::{
8    DelegatableNode, InvalidationKind, ModifierInvalidation, ModifierNode, ModifierNodeChain,
9    ModifierNodeElement, NodeCapabilities, NodeState,
10};
11
12#[derive(Clone)]
13struct ModifierLocalId(Rc<()>);
14
15impl ModifierLocalId {
16    fn new() -> Self {
17        Self(Rc::new(()))
18    }
19}
20
21impl fmt::Debug for ModifierLocalId {
22    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
23        f.debug_tuple("ModifierLocalId")
24            .field(&(Rc::as_ptr(&self.0) as usize))
25            .finish()
26    }
27}
28
29impl PartialEq for ModifierLocalId {
30    fn eq(&self, other: &Self) -> bool {
31        Rc::ptr_eq(&self.0, &other.0)
32    }
33}
34
35impl Eq for ModifierLocalId {}
36
37impl Hash for ModifierLocalId {
38    fn hash<H: Hasher>(&self, state: &mut H) {
39        Rc::as_ptr(&self.0).hash(state);
40    }
41}
42
43#[derive(Clone, Debug, PartialEq, Eq, Hash)]
44pub(crate) struct ModifierLocalToken {
45    id: ModifierLocalId,
46    type_id: TypeId,
47}
48
49impl ModifierLocalToken {
50    fn new(type_id: TypeId) -> Self {
51        let id = ModifierLocalId::new();
52        Self { id, type_id }
53    }
54
55    fn id(&self) -> ModifierLocalId {
56        self.id.clone()
57    }
58}
59
60/// Type-safe key referencing a modifier local value.
61#[derive(Clone)]
62pub struct ModifierLocalKey<T: 'static> {
63    token: ModifierLocalToken,
64    default: Rc<dyn Fn() -> T>,
65}
66
67impl<T: 'static> ModifierLocalKey<T> {
68    pub fn new(factory: impl Fn() -> T + 'static) -> Self {
69        Self {
70            token: ModifierLocalToken::new(TypeId::of::<T>()),
71            default: Rc::new(factory),
72        }
73    }
74
75    pub(crate) fn token(&self) -> ModifierLocalToken {
76        self.token.clone()
77    }
78
79    pub(crate) fn default_value(&self) -> T {
80        (self.default)()
81    }
82}
83
84impl<T: 'static> PartialEq for ModifierLocalKey<T> {
85    fn eq(&self, other: &Self) -> bool {
86        self.token == other.token
87    }
88}
89
90impl<T: 'static> Eq for ModifierLocalKey<T> {}
91
92impl<T: 'static> Hash for ModifierLocalKey<T> {
93    fn hash<H: Hasher>(&self, state: &mut H) {
94        self.token.hash(state);
95    }
96}
97
98/// Node responsible for providing a modifier local value.
99pub struct ModifierLocalProviderNode {
100    token: ModifierLocalToken,
101    value_factory: Rc<dyn Fn() -> Box<dyn Any>>,
102    value: Rc<dyn Any>,
103    version: u64,
104    state: NodeState,
105}
106
107impl ModifierLocalProviderNode {
108    fn new(token: ModifierLocalToken, factory: Rc<dyn Fn() -> Box<dyn Any>>) -> Self {
109        Self {
110            token,
111            value: Self::create_value(&factory),
112            value_factory: factory,
113            version: 0,
114            state: NodeState::new(),
115        }
116    }
117
118    fn update_value(&mut self) {
119        self.value = Self::create_value(&self.value_factory);
120        self.version = self.version.wrapping_add(1);
121    }
122
123    fn set_factory(&mut self, factory: Rc<dyn Fn() -> Box<dyn Any>>) {
124        self.value_factory = factory;
125        self.update_value();
126    }
127
128    fn token(&self) -> ModifierLocalToken {
129        self.token.clone()
130    }
131
132    fn value(&self) -> Rc<dyn Any> {
133        self.value.clone()
134    }
135
136    fn version(&self) -> u64 {
137        self.version
138    }
139
140    fn create_value(factory: &Rc<dyn Fn() -> Box<dyn Any>>) -> Rc<dyn Any> {
141        Rc::from(factory())
142    }
143}
144
145impl DelegatableNode for ModifierLocalProviderNode {
146    fn node_state(&self) -> &NodeState {
147        &self.state
148    }
149}
150
151impl ModifierNode for ModifierLocalProviderNode {}
152
153/// Node responsible for observing modifier local changes.
154pub struct ModifierLocalConsumerNode {
155    callback: Rc<dyn for<'a> Fn(&mut ModifierLocalReadScope<'a>)>,
156    state: NodeState,
157}
158
159impl ModifierLocalConsumerNode {
160    fn new(callback: Rc<dyn for<'a> Fn(&mut ModifierLocalReadScope<'a>)>) -> Self {
161        Self {
162            callback,
163            state: NodeState::new(),
164        }
165    }
166
167    fn notify(&self, scope: &mut ModifierLocalReadScope<'_>) {
168        (self.callback)(scope);
169    }
170
171    fn id(&self) -> usize {
172        self as *const Self as usize
173    }
174}
175
176#[derive(Clone)]
177pub(crate) struct ResolvedModifierLocal {
178    value: Rc<dyn Any>,
179    version: u64,
180    source: ModifierLocalSource,
181}
182
183impl ResolvedModifierLocal {
184    fn new(value: Rc<dyn Any>, version: u64, source: ModifierLocalSource) -> Self {
185        Self {
186            value,
187            version,
188            source,
189        }
190    }
191
192    pub(crate) fn value(&self) -> Rc<dyn Any> {
193        self.value.clone()
194    }
195
196    pub(crate) fn version(&self) -> u64 {
197        self.version
198    }
199
200    pub(crate) fn with_source(mut self, source: ModifierLocalSource) -> Self {
201        self.source = source;
202        self
203    }
204}
205
206#[derive(Clone, Copy, Debug, PartialEq, Eq)]
207pub(crate) enum ModifierLocalSource {
208    Chain,
209    Ancestor,
210}
211
212pub(crate) type ModifierLocalAncestorResolver<'a> =
213    dyn FnMut(&ModifierLocalToken) -> Option<ResolvedModifierLocal> + 'a;
214
215#[derive(Clone)]
216struct ProviderRecord {
217    value: Rc<dyn Any>,
218    version: u64,
219}
220
221impl ProviderRecord {
222    fn new(value: Rc<dyn Any>, version: u64) -> Self {
223        Self { value, version }
224    }
225
226    fn version(&self) -> u64 {
227        self.version
228    }
229
230    fn value(&self) -> &Rc<dyn Any> {
231        &self.value
232    }
233}
234
235#[derive(Clone)]
236struct DependencyRecord {
237    token: ModifierLocalToken,
238    source: DependencySource,
239    version: u64,
240}
241
242impl DependencyRecord {
243    fn from_chain(token: ModifierLocalToken, version: u64) -> Self {
244        Self {
245            token,
246            source: DependencySource::Chain,
247            version,
248        }
249    }
250
251    fn from_ancestor(token: ModifierLocalToken, version: u64) -> Self {
252        Self {
253            token,
254            source: DependencySource::Ancestor,
255            version,
256        }
257    }
258
259    fn from_default(token: ModifierLocalToken) -> Self {
260        Self {
261            token,
262            source: DependencySource::Default,
263            version: 0,
264        }
265    }
266
267    fn is_dirty(
268        &self,
269        providers: &HashMap<ModifierLocalId, ProviderRecord>,
270        ancestor_lookup: &mut ModifierLocalAncestorResolver<'_>,
271    ) -> bool {
272        match self.source {
273            DependencySource::Chain => match providers.get(&self.token.id()) {
274                Some(record) => record.version() != self.version,
275                None => true,
276            },
277            DependencySource::Ancestor => {
278                if providers.contains_key(&self.token.id()) {
279                    return true;
280                }
281                ancestor_lookup(&self.token)
282                    .map(|resolved| resolved.version() != self.version)
283                    .unwrap_or(true)
284            }
285            DependencySource::Default => {
286                providers.contains_key(&self.token.id()) || ancestor_lookup(&self.token).is_some()
287            }
288        }
289    }
290}
291
292#[derive(Clone, Copy)]
293enum DependencySource {
294    Chain,
295    Ancestor,
296    Default,
297}
298
299struct ConsumerState {
300    dependencies: Vec<DependencyRecord>,
301}
302
303impl ConsumerState {
304    fn new(dependencies: Vec<DependencyRecord>) -> Self {
305        Self { dependencies }
306    }
307
308    fn needs_update(
309        &self,
310        providers: &HashMap<ModifierLocalId, ProviderRecord>,
311        ancestor_lookup: &mut ModifierLocalAncestorResolver<'_>,
312    ) -> bool {
313        if self.dependencies.is_empty() {
314            return true;
315        }
316        self.dependencies
317            .iter()
318            .any(|dependency| dependency.is_dirty(providers, ancestor_lookup))
319    }
320}
321
322impl DelegatableNode for ModifierLocalConsumerNode {
323    fn node_state(&self) -> &NodeState {
324        &self.state
325    }
326}
327
328impl ModifierNode for ModifierLocalConsumerNode {}
329
330#[derive(Clone)]
331pub struct ModifierLocalProviderElement {
332    token: ModifierLocalToken,
333    factory: Rc<dyn Fn() -> Box<dyn Any>>,
334}
335
336impl ModifierLocalProviderElement {
337    pub fn new<T, F>(key: ModifierLocalKey<T>, factory: F) -> Self
338    where
339        T: 'static,
340        F: Fn() -> T + 'static,
341    {
342        let erased = Rc::new(move || -> Box<dyn Any> { Box::new(factory()) });
343        Self {
344            token: key.token(),
345            factory: erased,
346        }
347    }
348}
349
350impl fmt::Debug for ModifierLocalProviderElement {
351    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
352        f.debug_struct("ModifierLocalProviderElement")
353            .field("id", &self.token.id())
354            .finish()
355    }
356}
357
358impl PartialEq for ModifierLocalProviderElement {
359    fn eq(&self, other: &Self) -> bool {
360        // Type-based matching: compare only tokens, not factory closures
361        // Nodes are updated via update() method, preserving behavior
362        self.token == other.token
363    }
364}
365
366impl Eq for ModifierLocalProviderElement {}
367
368impl Hash for ModifierLocalProviderElement {
369    fn hash<H: Hasher>(&self, state: &mut H) {
370        // Consistent hash based on token only
371        "modifier_local_provider".hash(state);
372        self.token.hash(state);
373    }
374}
375
376impl ModifierNodeElement for ModifierLocalProviderElement {
377    type Node = ModifierLocalProviderNode;
378
379    fn create(&self) -> Self::Node {
380        ModifierLocalProviderNode::new(self.token.clone(), self.factory.clone())
381    }
382
383    fn update(&self, node: &mut Self::Node) {
384        node.set_factory(self.factory.clone());
385    }
386
387    fn capabilities(&self) -> NodeCapabilities {
388        NodeCapabilities::MODIFIER_LOCALS
389    }
390
391    fn always_update(&self) -> bool {
392        // Factory closure might change even if token is same
393        true
394    }
395}
396
397#[derive(Clone)]
398pub struct ModifierLocalConsumerElement {
399    callback: Rc<dyn for<'a> Fn(&mut ModifierLocalReadScope<'a>)>,
400}
401
402impl ModifierLocalConsumerElement {
403    pub fn new<F>(callback: F) -> Self
404    where
405        F: for<'a> Fn(&mut ModifierLocalReadScope<'a>) + 'static,
406    {
407        Self {
408            callback: Rc::new(callback),
409        }
410    }
411}
412
413impl fmt::Debug for ModifierLocalConsumerElement {
414    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
415        f.write_str("ModifierLocalConsumerElement")
416    }
417}
418
419impl PartialEq for ModifierLocalConsumerElement {
420    fn eq(&self, _other: &Self) -> bool {
421        // Type-based matching: always equal for same type
422        // Nodes are updated via update() method, preserving behavior
423        true
424    }
425}
426
427impl Eq for ModifierLocalConsumerElement {}
428
429impl Hash for ModifierLocalConsumerElement {
430    fn hash<H: Hasher>(&self, state: &mut H) {
431        // Consistent hash for type-based matching
432        "modifier_local_consumer".hash(state);
433    }
434}
435
436impl ModifierNodeElement for ModifierLocalConsumerElement {
437    type Node = ModifierLocalConsumerNode;
438
439    fn create(&self) -> Self::Node {
440        ModifierLocalConsumerNode::new(self.callback.clone())
441    }
442
443    fn update(&self, node: &mut Self::Node) {
444        node.callback = self.callback.clone();
445    }
446
447    fn capabilities(&self) -> NodeCapabilities {
448        NodeCapabilities::MODIFIER_LOCALS
449    }
450
451    fn always_update(&self) -> bool {
452        // Callback closure might change
453        true
454    }
455}
456
457/// Lightweight read scope surfaced to modifier local consumers.
458pub struct ModifierLocalReadScope<'a> {
459    providers: &'a HashMap<ModifierLocalId, ProviderRecord>,
460    ancestor_lookup: &'a mut ModifierLocalAncestorResolver<'a>,
461    dependencies: &'a mut Vec<DependencyRecord>,
462    fallbacks: HashMap<ModifierLocalId, Rc<dyn Any>>,
463}
464
465impl<'a> ModifierLocalReadScope<'a> {
466    fn new(
467        providers: &'a HashMap<ModifierLocalId, ProviderRecord>,
468        ancestor_lookup: &'a mut ModifierLocalAncestorResolver<'a>,
469        dependencies: &'a mut Vec<DependencyRecord>,
470    ) -> Self {
471        Self {
472            providers,
473            ancestor_lookup,
474            dependencies,
475            fallbacks: HashMap::new(),
476        }
477    }
478
479    pub fn get<T: 'static>(&mut self, key: &ModifierLocalKey<T>) -> &T {
480        let token = key.token();
481        if let Some(record) = self.providers.get(&token.id()) {
482            self.dependencies
483                .push(DependencyRecord::from_chain(token, record.version()));
484            if let Some(value) = record.value().downcast_ref::<T>() {
485                return value;
486            }
487            return self.default_value_for(key);
488        }
489
490        if let Some(resolved) = (self.ancestor_lookup)(&token) {
491            self.dependencies.push(DependencyRecord::from_ancestor(
492                token.clone(),
493                resolved.version(),
494            ));
495            let resolved_value = resolved.value();
496            if resolved_value.downcast_ref::<T>().is_some() {
497                self.fallbacks.insert(token.id(), resolved_value);
498                return self.fallback_value_for(key);
499            }
500            return self.default_value_for(key);
501        }
502
503        self.dependencies
504            .push(DependencyRecord::from_default(token));
505        self.default_value_for(key)
506    }
507
508    fn fallback_value_for<T: 'static>(&mut self, key: &ModifierLocalKey<T>) -> &T {
509        let id = key.token().id();
510        if self
511            .fallbacks
512            .get(&id)
513            .and_then(|value| value.downcast_ref::<T>())
514            .is_none()
515        {
516            self.fallbacks
517                .insert(id.clone(), Rc::new(key.default_value()) as Rc<dyn Any>);
518        }
519
520        match self
521            .fallbacks
522            .get(&id)
523            .and_then(|value| value.downcast_ref::<T>())
524        {
525            Some(value) => value,
526            None => Box::leak(Box::new(key.default_value())),
527        }
528    }
529
530    fn default_value_for<T: 'static>(&mut self, key: &ModifierLocalKey<T>) -> &T {
531        self.fallback_value_for(key)
532    }
533}
534
535#[cfg(test)]
536mod tests {
537    use super::*;
538
539    fn empty_ancestor_lookup(_: &ModifierLocalToken) -> Option<ResolvedModifierLocal> {
540        None
541    }
542
543    #[test]
544    fn modifier_local_keys_do_not_use_process_global_counter() {
545        let source = include_str!("local.rs");
546        assert!(!source.contains(concat!("NEXT_", "MODIFIER_LOCAL_ID")));
547        assert!(!source.contains(concat!("Atomic", "U64")));
548    }
549
550    #[test]
551    fn modifier_local_key_identity_is_retained_per_key() {
552        let first = ModifierLocalKey::new(|| 1_i32);
553        let first_clone = first.clone();
554        let second = ModifierLocalKey::new(|| 1_i32);
555
556        assert_eq!(first.token(), first_clone.token());
557        assert_ne!(first.token(), second.token());
558    }
559
560    #[test]
561    fn modifier_local_provider_type_mismatch_falls_back_to_default() {
562        let key = ModifierLocalKey::new(|| 7_i32);
563        let mut providers = HashMap::new();
564        providers.insert(
565            key.token().id(),
566            ProviderRecord::new(Rc::new(String::from("wrong")) as Rc<dyn Any>, 4),
567        );
568        let mut dependencies = Vec::new();
569        let mut ancestor_lookup = empty_ancestor_lookup;
570        let mut scope =
571            ModifierLocalReadScope::new(&providers, &mut ancestor_lookup, &mut dependencies);
572
573        assert_eq!(*scope.get(&key), 7);
574        assert!(matches!(dependencies[0].source, DependencySource::Chain));
575        assert_eq!(dependencies[0].version, 4);
576    }
577
578    #[test]
579    fn modifier_local_ancestor_type_mismatch_falls_back_to_default() {
580        let key = ModifierLocalKey::new(|| 11_i32);
581        let providers = HashMap::new();
582        let mut dependencies = Vec::new();
583        let mut ancestor_lookup = |token: &ModifierLocalToken| {
584            if *token == key.token() {
585                Some(ResolvedModifierLocal::new(
586                    Rc::new(String::from("wrong")) as Rc<dyn Any>,
587                    9,
588                    ModifierLocalSource::Ancestor,
589                ))
590            } else {
591                None
592            }
593        };
594        let mut scope =
595            ModifierLocalReadScope::new(&providers, &mut ancestor_lookup, &mut dependencies);
596
597        assert_eq!(*scope.get(&key), 11);
598        assert!(matches!(dependencies[0].source, DependencySource::Ancestor));
599        assert_eq!(dependencies[0].version, 9);
600    }
601
602    #[test]
603    fn modifier_local_fallback_cache_reconciles_mismatched_cached_value() {
604        let key = ModifierLocalKey::new(|| 23_i32);
605        let providers = HashMap::new();
606        let mut dependencies = Vec::new();
607        let mut ancestor_lookup = empty_ancestor_lookup;
608        let cache_is_typed = {
609            let mut scope =
610                ModifierLocalReadScope::new(&providers, &mut ancestor_lookup, &mut dependencies);
611            scope.fallbacks.insert(
612                key.token().id(),
613                Rc::new(String::from("wrong")) as Rc<dyn Any>,
614            );
615
616            assert_eq!(*scope.get(&key), 23);
617            assert_eq!(*scope.get(&key), 23);
618            scope
619                .fallbacks
620                .get(&key.token().id())
621                .and_then(|value| value.downcast_ref::<i32>())
622                .is_some()
623        };
624
625        assert_eq!(dependencies.len(), 2);
626        assert!(cache_is_typed);
627    }
628}
629
630#[derive(Default)]
631pub struct ModifierLocalManager {
632    providers: HashMap<ModifierLocalId, ProviderRecord>,
633    consumers: HashMap<usize, ConsumerState>,
634}
635
636impl ModifierLocalManager {
637    pub fn new() -> Self {
638        Self::default()
639    }
640
641    #[allow(private_interfaces)]
642    pub fn sync(
643        &mut self,
644        chain: &ModifierNodeChain,
645        ancestor_lookup: &mut ModifierLocalAncestorResolver<'_>,
646    ) -> Vec<ModifierInvalidation> {
647        if !chain.has_capability(NodeCapabilities::MODIFIER_LOCALS) {
648            self.providers.clear();
649            self.consumers.clear();
650            return Vec::new();
651        }
652
653        let mut providers: HashMap<ModifierLocalId, ProviderRecord> = HashMap::new();
654        let mut seen_consumers = HashSet::new();
655        let mut invalidations = Vec::new();
656
657        chain.for_each_node_with_capability(NodeCapabilities::MODIFIER_LOCALS, |_ref, node| {
658            if let Some(provider) = node.as_any().downcast_ref::<ModifierLocalProviderNode>() {
659                providers.insert(
660                    provider.token().id(),
661                    ProviderRecord::new(provider.value(), provider.version()),
662                );
663                return;
664            }
665
666            if let Some(consumer) = node.as_any().downcast_ref::<ModifierLocalConsumerNode>() {
667                let id = consumer.id();
668                seen_consumers.insert(id);
669                let needs_update = self
670                    .consumers
671                    .get(&id)
672                    .map(|state| state.needs_update(&providers, ancestor_lookup))
673                    .unwrap_or(true);
674                if !needs_update {
675                    return;
676                }
677
678                let mut dependencies = Vec::new();
679                {
680                    let mut scope =
681                        ModifierLocalReadScope::new(&providers, ancestor_lookup, &mut dependencies);
682                    consumer.notify(&mut scope);
683                }
684                self.consumers.insert(id, ConsumerState::new(dependencies));
685                if !invalidations
686                    .iter()
687                    .any(|entry: &ModifierInvalidation| entry.kind() == InvalidationKind::Layout)
688                {
689                    invalidations.push(ModifierInvalidation::new(
690                        InvalidationKind::Layout,
691                        NodeCapabilities::LAYOUT | NodeCapabilities::MODIFIER_LOCALS,
692                    ));
693                }
694            }
695        });
696
697        self.providers = providers;
698        self.consumers.retain(|id, _| seen_consumers.contains(id));
699
700        invalidations
701    }
702
703    pub(crate) fn resolve(&self, token: &ModifierLocalToken) -> Option<ResolvedModifierLocal> {
704        self.providers.get(&token.id()).map(|record| {
705            ResolvedModifierLocal::new(
706                record.value().clone(),
707                record.version(),
708                ModifierLocalSource::Chain,
709            )
710        })
711    }
712}