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#[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
98pub 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
153pub 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
363 }
364}
365
366impl Eq for ModifierLocalProviderElement {}
367
368impl Hash for ModifierLocalProviderElement {
369 fn hash<H: Hasher>(&self, state: &mut H) {
370 "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 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 true
424 }
425}
426
427impl Eq for ModifierLocalConsumerElement {}
428
429impl Hash for ModifierLocalConsumerElement {
430 fn hash<H: Hasher>(&self, state: &mut H) {
431 "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 true
454 }
455}
456
457pub 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}