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 {
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
449pub 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}