Skip to main content

cranpose_ui/modifier/
local.rs

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