1use super::{
2 COMPUTING, TrackingGuard,
3 graph::{Observer, ObserverKind, Source, track},
4 untrack,
5};
6use std::{
7 cell::{Cell, RefCell},
8 rc::Rc,
9};
10
11#[must_use = "memos are lazy; retain a handle and read it to use the computation"]
20pub struct Memo<T>(Rc<MemoInner<T>>);
21
22impl<T> Clone for Memo<T> {
23 fn clone(&self) -> Self {
24 Self(self.0.clone())
25 }
26}
27
28pub fn memo<T: PartialEq + 'static>(compute: impl Fn() -> T + 'static) -> Memo<T> {
30 memo_with_eq(compute, T::eq)
31}
32
33pub fn memo_with_eq<T: 'static>(
36 compute: impl Fn() -> T + 'static,
37 equal: impl Fn(&T, &T) -> bool + 'static,
38) -> Memo<T> {
39 Memo(Rc::<MemoInner<T>>::new_cyclic(|weak| MemoInner {
40 source: Rc::new(Source::new(Some(weak.clone()))),
41 observer: Observer::new(ObserverKind::Memo(weak.clone())),
42 value: RefCell::new(None),
43 compute: Box::new(compute),
44 equal: Box::new(equal),
45 stale: Cell::new(true),
46 needs_compute: Cell::new(true),
47 running: Cell::new(false),
48 }))
49}
50
51type Equal<T> = dyn Fn(&T, &T) -> bool;
52
53struct MemoInner<T> {
54 source: Rc<Source>,
55 observer: Rc<Observer>,
56 value: RefCell<Option<T>>,
57 compute: Box<dyn Fn() -> T>,
58 equal: Box<Equal<T>>,
59 stale: Cell<bool>,
60 needs_compute: Cell<bool>,
61 running: Cell<bool>,
62}
63
64impl<T> Drop for MemoInner<T> {
65 fn drop(&mut self) {
66 self.observer.unsubscribe();
67 untrack(|| drop(self.value.get_mut().take()));
70 }
71}
72
73pub(super) trait MemoNode {
74 fn refresh(&self);
75 fn invalidate(&self);
76 fn source(&self) -> &Source;
77}
78
79struct EvaluationGuard<'a>(&'a Cell<bool>);
80
81impl Drop for EvaluationGuard<'_> {
82 fn drop(&mut self) {
83 self.0.set(false);
84 COMPUTING.with(|depth| depth.set(depth.get() - 1));
85 }
86}
87
88impl<T: 'static> MemoNode for MemoInner<T> {
89 fn refresh(&self) {
90 assert!(
91 !self.running.get(),
92 "reactive cycle: a memo depends on itself"
93 );
94 if !self.stale.get() {
95 return;
96 }
97 self.running.set(true);
98 COMPUTING.with(|depth| depth.set(depth.get() + 1));
99 let _evaluation = EvaluationGuard(&self.running);
100 if !self.needs_compute.get() && !untrack(|| self.observer.changed()) {
101 self.stale.set(false);
102 return;
103 }
104 self.needs_compute.set(true);
107 self.observer.unsubscribe();
108 let next = {
109 let _tracking = TrackingGuard::replace(Some(Rc::downgrade(&self.observer)));
110 (self.compute)()
111 };
112 let equal = untrack(|| {
113 self.value
114 .borrow()
115 .as_ref()
116 .is_some_and(|old| (self.equal)(old, &next))
117 });
118 if equal {
119 untrack(|| drop(next));
120 } else {
121 let old = self.value.replace(Some(next));
122 self.source.advance();
123 untrack(|| drop(old));
125 }
126 self.needs_compute.set(false);
127 self.stale.set(false);
128 }
129
130 fn invalidate(&self) {
131 self.stale.set(true);
132 }
133 fn source(&self) -> &Source {
134 &self.source
135 }
136}
137
138impl<T: 'static> Memo<T> {
139 pub fn with<R>(&self, read: impl FnOnce(&T) -> R) -> R {
142 self.0.refresh();
143 track(&self.0.source);
144 read(self.0.value.borrow().as_ref().expect("memo evaluated"))
145 }
146
147 pub fn with_untracked<R>(&self, read: impl FnOnce(&T) -> R) -> R {
149 untrack(|| self.with(read))
150 }
151}
152
153impl<T: Clone + 'static> Memo<T> {
154 pub fn get(&self) -> T {
155 self.with(Clone::clone)
156 }
157 pub fn get_untracked(&self) -> T {
158 self.with_untracked(Clone::clone)
159 }
160}