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
100pub struct ModifierLocalProviderNode {
101    token: ModifierLocalToken,
102    value_factory: Rc<dyn Fn() -> Box<dyn Any>>,
103    value: Rc<dyn Any>,
104    version: u64,
105    state: NodeState,
106}
107
108impl ModifierLocalProviderNode {
109    fn new(token: ModifierLocalToken, factory: Rc<dyn Fn() -> Box<dyn Any>>) -> Self {
110        Self {
111            token,
112            value: Self::create_value(&factory),
113            value_factory: factory,
114            version: 0,
115            state: NodeState::new(),
116        }
117    }
118
119    fn update_value(&mut self) {
120        self.value = Self::create_value(&self.value_factory);
121        self.version = self.version.wrapping_add(1);
122    }
123
124    fn set_factory(&mut self, factory: Rc<dyn Fn() -> Box<dyn Any>>) {
125        self.value_factory = factory;
126        self.update_value();
127    }
128
129    fn token(&self) -> ModifierLocalToken {
130        self.token.clone()
131    }
132
133    fn value(&self) -> Rc<dyn Any> {
134        self.value.clone()
135    }
136
137    fn version(&self) -> u64 {
138        self.version
139    }
140
141    fn create_value(factory: &Rc<dyn Fn() -> Box<dyn Any>>) -> Rc<dyn Any> {
142        Rc::from(factory())
143    }
144}
145
146impl DelegatableNode for ModifierLocalProviderNode {
147    fn node_state(&self) -> &NodeState {
148        &self.state
149    }
150}
151
152impl ModifierNode for ModifierLocalProviderNode {}
153
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        self.token == other.token
361    }
362}
363
364impl Eq for ModifierLocalProviderElement {}
365
366impl Hash for ModifierLocalProviderElement {
367    fn hash<H: Hasher>(&self, state: &mut H) {
368        "modifier_local_provider".hash(state);
369        self.token.hash(state);
370    }
371}
372
373impl ModifierNodeElement for ModifierLocalProviderElement {
374    type Node = ModifierLocalProviderNode;
375
376    fn create(&self) -> Self::Node {
377        ModifierLocalProviderNode::new(self.token.clone(), self.factory.clone())
378    }
379
380    fn update(&self, node: &mut Self::Node) {
381        node.set_factory(self.factory.clone());
382    }
383
384    fn capabilities(&self) -> NodeCapabilities {
385        NodeCapabilities::MODIFIER_LOCALS
386    }
387
388    fn always_update(&self) -> bool {
389        true
390    }
391}
392
393#[derive(Clone)]
394pub struct ModifierLocalConsumerElement {
395    callback: Rc<dyn for<'a> Fn(&mut ModifierLocalReadScope<'a>)>,
396}
397
398impl ModifierLocalConsumerElement {
399    pub fn new<F>(callback: F) -> Self
400    where
401        F: for<'a> Fn(&mut ModifierLocalReadScope<'a>) + 'static,
402    {
403        Self {
404            callback: Rc::new(callback),
405        }
406    }
407}
408
409impl fmt::Debug for ModifierLocalConsumerElement {
410    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
411        f.write_str("ModifierLocalConsumerElement")
412    }
413}
414
415impl PartialEq for ModifierLocalConsumerElement {
416    fn eq(&self, _other: &Self) -> bool {
417        true
418    }
419}
420
421impl Eq for ModifierLocalConsumerElement {}
422
423impl Hash for ModifierLocalConsumerElement {
424    fn hash<H: Hasher>(&self, state: &mut H) {
425        "modifier_local_consumer".hash(state);
426    }
427}
428
429impl ModifierNodeElement for ModifierLocalConsumerElement {
430    type Node = ModifierLocalConsumerNode;
431
432    fn create(&self) -> Self::Node {
433        ModifierLocalConsumerNode::new(self.callback.clone())
434    }
435
436    fn update(&self, node: &mut Self::Node) {
437        node.callback = self.callback.clone();
438    }
439
440    fn capabilities(&self) -> NodeCapabilities {
441        NodeCapabilities::MODIFIER_LOCALS
442    }
443
444    fn always_update(&self) -> bool {
445        true
446    }
447}
448
449/// Lightweight read scope surfaced to modifier local consumers.
450pub struct ModifierLocalReadScope<'a> {
451    providers: &'a HashMap<ModifierLocalId, ProviderRecord>,
452    ancestor_lookup: &'a mut ModifierLocalAncestorResolver<'a>,
453    dependencies: &'a mut Vec<DependencyRecord>,
454    fallbacks: HashMap<ModifierLocalId, Rc<dyn Any>>,
455}
456
457impl<'a> ModifierLocalReadScope<'a> {
458    fn new(
459        providers: &'a HashMap<ModifierLocalId, ProviderRecord>,
460        ancestor_lookup: &'a mut ModifierLocalAncestorResolver<'a>,
461        dependencies: &'a mut Vec<DependencyRecord>,
462    ) -> Self {
463        Self {
464            providers,
465            ancestor_lookup,
466            dependencies,
467            fallbacks: HashMap::new(),
468        }
469    }
470
471    pub fn get<T: 'static>(&mut self, key: &ModifierLocalKey<T>) -> &T {
472        let token = key.token();
473        if let Some(record) = self.providers.get(&token.id()) {
474            self.dependencies
475                .push(DependencyRecord::from_chain(token, record.version()));
476            if let Some(value) = record.value().downcast_ref::<T>() {
477                return value;
478            }
479            return self.default_value_for(key);
480        }
481
482        if let Some(resolved) = (self.ancestor_lookup)(&token) {
483            self.dependencies.push(DependencyRecord::from_ancestor(
484                token.clone(),
485                resolved.version(),
486            ));
487            let resolved_value = resolved.value();
488            if resolved_value.downcast_ref::<T>().is_some() {
489                self.fallbacks.insert(token.id(), resolved_value);
490                return self.fallback_value_for(key);
491            }
492            return self.default_value_for(key);
493        }
494
495        self.dependencies
496            .push(DependencyRecord::from_default(token));
497        self.default_value_for(key)
498    }
499
500    fn fallback_value_for<T: 'static>(&mut self, key: &ModifierLocalKey<T>) -> &T {
501        let id = key.token().id();
502        if self
503            .fallbacks
504            .get(&id)
505            .and_then(|value| value.downcast_ref::<T>())
506            .is_none()
507        {
508            self.fallbacks
509                .insert(id.clone(), Rc::new(key.default_value()) as Rc<dyn Any>);
510        }
511
512        match self
513            .fallbacks
514            .get(&id)
515            .and_then(|value| value.downcast_ref::<T>())
516        {
517            Some(value) => value,
518            None => Box::leak(Box::new(key.default_value())),
519        }
520    }
521
522    fn default_value_for<T: 'static>(&mut self, key: &ModifierLocalKey<T>) -> &T {
523        self.fallback_value_for(key)
524    }
525}
526
527#[cfg(test)]
528mod tests {
529    use super::*;
530
531    fn empty_ancestor_lookup(_: &ModifierLocalToken) -> Option<ResolvedModifierLocal> {
532        None
533    }
534
535    #[test]
536    fn modifier_local_keys_do_not_use_process_global_counter() {
537        let source = include_str!("local.rs");
538        assert!(!source.contains(concat!("NEXT_", "MODIFIER_LOCAL_ID")));
539        assert!(!source.contains(concat!("Atomic", "U64")));
540    }
541
542    #[test]
543    fn modifier_local_key_identity_is_retained_per_key() {
544        let first = ModifierLocalKey::new(|| 1_i32);
545        let first_clone = first.clone();
546        let second = ModifierLocalKey::new(|| 1_i32);
547
548        assert_eq!(first.token(), first_clone.token());
549        assert_ne!(first.token(), second.token());
550    }
551
552    #[test]
553    fn modifier_local_provider_type_mismatch_falls_back_to_default() {
554        let key = ModifierLocalKey::new(|| 7_i32);
555        let mut providers = HashMap::new();
556        providers.insert(
557            key.token().id(),
558            ProviderRecord::new(Rc::new(String::from("wrong")) as Rc<dyn Any>, 4),
559        );
560        let mut dependencies = Vec::new();
561        let mut ancestor_lookup = empty_ancestor_lookup;
562        let mut scope =
563            ModifierLocalReadScope::new(&providers, &mut ancestor_lookup, &mut dependencies);
564
565        assert_eq!(*scope.get(&key), 7);
566        assert!(matches!(dependencies[0].source, DependencySource::Chain));
567        assert_eq!(dependencies[0].version, 4);
568    }
569
570    #[test]
571    fn modifier_local_ancestor_type_mismatch_falls_back_to_default() {
572        let key = ModifierLocalKey::new(|| 11_i32);
573        let providers = HashMap::new();
574        let mut dependencies = Vec::new();
575        let mut ancestor_lookup = |token: &ModifierLocalToken| {
576            if *token == key.token() {
577                Some(ResolvedModifierLocal::new(
578                    Rc::new(String::from("wrong")) as Rc<dyn Any>,
579                    9,
580                    ModifierLocalSource::Ancestor,
581                ))
582            } else {
583                None
584            }
585        };
586        let mut scope =
587            ModifierLocalReadScope::new(&providers, &mut ancestor_lookup, &mut dependencies);
588
589        assert_eq!(*scope.get(&key), 11);
590        assert!(matches!(dependencies[0].source, DependencySource::Ancestor));
591        assert_eq!(dependencies[0].version, 9);
592    }
593
594    #[test]
595    fn modifier_local_fallback_cache_reconciles_mismatched_cached_value() {
596        let key = ModifierLocalKey::new(|| 23_i32);
597        let providers = HashMap::new();
598        let mut dependencies = Vec::new();
599        let mut ancestor_lookup = empty_ancestor_lookup;
600        let cache_is_typed = {
601            let mut scope =
602                ModifierLocalReadScope::new(&providers, &mut ancestor_lookup, &mut dependencies);
603            scope.fallbacks.insert(
604                key.token().id(),
605                Rc::new(String::from("wrong")) as Rc<dyn Any>,
606            );
607
608            assert_eq!(*scope.get(&key), 23);
609            assert_eq!(*scope.get(&key), 23);
610            scope
611                .fallbacks
612                .get(&key.token().id())
613                .and_then(|value| value.downcast_ref::<i32>())
614                .is_some()
615        };
616
617        assert_eq!(dependencies.len(), 2);
618        assert!(cache_is_typed);
619    }
620}
621
622#[derive(Default)]
623pub struct ModifierLocalManager {
624    providers: HashMap<ModifierLocalId, ProviderRecord>,
625    consumers: HashMap<usize, ConsumerState>,
626}
627
628impl ModifierLocalManager {
629    pub fn new() -> Self {
630        Self::default()
631    }
632
633    #[allow(private_interfaces)]
634    pub fn sync(
635        &mut self,
636        chain: &ModifierNodeChain,
637        ancestor_lookup: &mut ModifierLocalAncestorResolver<'_>,
638    ) -> Vec<ModifierInvalidation> {
639        if !chain.has_capability(NodeCapabilities::MODIFIER_LOCALS) {
640            self.providers.clear();
641            self.consumers.clear();
642            return Vec::new();
643        }
644
645        let mut providers: HashMap<ModifierLocalId, ProviderRecord> = HashMap::new();
646        let mut seen_consumers = HashSet::new();
647        let mut invalidations = Vec::new();
648
649        chain.for_each_node_with_capability(NodeCapabilities::MODIFIER_LOCALS, |_ref, node| {
650            if let Some(provider) = node.as_any().downcast_ref::<ModifierLocalProviderNode>() {
651                providers.insert(
652                    provider.token().id(),
653                    ProviderRecord::new(provider.value(), provider.version()),
654                );
655                return;
656            }
657
658            if let Some(consumer) = node.as_any().downcast_ref::<ModifierLocalConsumerNode>() {
659                let id = consumer.id();
660                seen_consumers.insert(id);
661                let needs_update = self
662                    .consumers
663                    .get(&id)
664                    .map(|state| state.needs_update(&providers, ancestor_lookup))
665                    .unwrap_or(true);
666                if !needs_update {
667                    return;
668                }
669
670                let mut dependencies = Vec::new();
671                {
672                    let mut scope =
673                        ModifierLocalReadScope::new(&providers, ancestor_lookup, &mut dependencies);
674                    consumer.notify(&mut scope);
675                }
676                self.consumers.insert(id, ConsumerState::new(dependencies));
677                if !invalidations
678                    .iter()
679                    .any(|entry: &ModifierInvalidation| entry.kind() == InvalidationKind::Layout)
680                {
681                    invalidations.push(ModifierInvalidation::new(
682                        InvalidationKind::Layout,
683                        NodeCapabilities::LAYOUT | NodeCapabilities::MODIFIER_LOCALS,
684                    ));
685                }
686            }
687        });
688
689        self.providers = providers;
690        self.consumers.retain(|id, _| seen_consumers.contains(id));
691
692        invalidations
693    }
694
695    pub(crate) fn resolve(&self, token: &ModifierLocalToken) -> Option<ResolvedModifierLocal> {
696        self.providers.get(&token.id()).map(|record| {
697            ResolvedModifierLocal::new(
698                record.value().clone(),
699                record.version(),
700                ModifierLocalSource::Chain,
701            )
702        })
703    }
704}