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#[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 {
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
155pub 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 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 "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 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 true
426 }
427}
428
429impl Eq for ModifierLocalConsumerElement {}
430
431impl Hash for ModifierLocalConsumerElement {
432 fn hash<H: Hasher>(&self, state: &mut H) {
433 "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 true
456 }
457}
458
459pub 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}