1use super::{
2 COMPUTING, RENDERING, RenderGuard, TrackingGuard,
3 graph::{Observer, ObserverKind, Source, track},
4 untrack,
5};
6use std::{
7 cell::{Cell, OnceCell, RefCell},
8 rc::{Rc, Weak},
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 graph: OnceCell::new(),
41 weak_self: weak.clone(),
42 observer_id: Observer::reserve_id(),
43 value: RefCell::new(None),
44 compute: Box::new(compute),
45 equal: Box::new(equal),
46 stale: Cell::new(true),
47 needs_compute: Cell::new(true),
48 running: Cell::new(false),
49 }))
50}
51
52type Equal<T> = dyn Fn(&T, &T) -> bool;
53
54struct MemoGraph {
55 source: Rc<Source>,
56 observer: Rc<Observer>,
57}
58
59struct MemoInner<T> {
60 graph: OnceCell<MemoGraph>,
61 weak_self: Weak<MemoInner<T>>,
62 observer_id: u64,
63 value: RefCell<Option<T>>,
64 compute: Box<dyn Fn() -> T>,
65 equal: Box<Equal<T>>,
66 stale: Cell<bool>,
67 needs_compute: Cell<bool>,
68 running: Cell<bool>,
69}
70
71impl<T> Drop for MemoInner<T> {
72 fn drop(&mut self) {
73 if let Some(graph) = self.graph.get() {
74 graph.observer.unsubscribe();
75 }
76 if let Some(value) = self.value.get_mut().take() {
79 untrack(|| drop(value));
80 }
81 }
82}
83
84impl<T: 'static> MemoInner<T> {
85 fn graph(&self) -> &MemoGraph {
86 self.graph.get_or_init(|| MemoGraph {
88 source: Rc::new(Source::new(Some(self.weak_self.clone()))),
89 observer: Observer::with_id(
90 self.observer_id,
91 ObserverKind::Memo(self.weak_self.clone()),
92 ),
93 })
94 }
95}
96
97pub(super) trait MemoNode {
98 fn refresh(&self);
99 fn invalidate(&self) -> &Source;
100}
101
102struct EvaluationGuard<'a>(&'a Cell<bool>);
103
104impl Drop for EvaluationGuard<'_> {
105 fn drop(&mut self) {
106 self.0.set(false);
107 COMPUTING.with(|depth| depth.set(depth.get() - 1));
108 }
109}
110
111impl<T: 'static> MemoNode for MemoInner<T> {
112 fn refresh(&self) {
113 let _committed = RenderGuard::replace(false);
115 assert!(
116 !self.running.get(),
117 "reactive cycle: a memo depends on itself"
118 );
119 if !self.stale.get() {
120 return;
121 }
122 self.running.set(true);
123 COMPUTING.with(|depth| depth.set(depth.get() + 1));
124 let _evaluation = EvaluationGuard(&self.running);
125 let graph = self.graph();
126 if !self.needs_compute.get() && !untrack(|| graph.observer.changed()) {
127 self.stale.set(false);
128 return;
129 }
130 self.needs_compute.set(true);
133 let next = {
134 let _run = graph.observer.begin();
135 let _tracking = TrackingGuard::replace(Some(Rc::downgrade(&graph.observer)));
136 (self.compute)()
137 };
138 let equal = untrack(|| {
139 self.value
140 .borrow()
141 .as_ref()
142 .is_some_and(|old| (self.equal)(old, &next))
143 });
144 if equal {
145 untrack(|| drop(next));
146 } else {
147 let old = self.value.replace(Some(next));
148 graph.source.advance();
149 untrack(|| drop(old));
151 }
152 self.needs_compute.set(false);
153 self.stale.set(false);
154 }
155
156 fn invalidate(&self) -> &Source {
157 self.stale.set(true);
158 &self.graph.get().expect("memo graph initialized").source
160 }
161}
162
163impl<T: 'static> Memo<T> {
164 pub fn with<R>(&self, read: impl FnOnce(&T) -> R) -> R {
167 if RENDERING.with(Cell::get) {
168 return self.with_candidate(read);
169 }
170 self.0.refresh();
171 track(&self.0.graph.get().expect("memo graph initialized").source);
173 read(self.0.value.borrow().as_ref().expect("memo evaluated"))
174 }
175
176 fn with_candidate<R>(&self, read: impl FnOnce(&T) -> R) -> R {
177 struct Candidate<T>(Option<T>);
178 impl<T> Drop for Candidate<T> {
179 fn drop(&mut self) {
180 untrack(|| drop(self.0.take()));
181 }
182 }
183 let value = {
184 assert!(
185 !self.0.running.replace(true),
186 "reactive cycle: a memo depends on itself"
187 );
188 COMPUTING.with(|depth| depth.set(depth.get() + 1));
189 let _evaluation = EvaluationGuard(&self.0.running);
190 Candidate(Some((self.0.compute)()))
193 };
194 read(value.0.as_ref().expect("candidate evaluated"))
195 }
196
197 pub fn with_untracked<R>(&self, read: impl FnOnce(&T) -> R) -> R {
199 untrack(|| self.with(read))
200 }
201}
202
203impl<T: Clone + 'static> Memo<T> {
204 pub fn get(&self) -> T {
205 self.with(Clone::clone)
206 }
207 pub fn get_untracked(&self) -> T {
208 self.with_untracked(Clone::clone)
209 }
210}