1use super::{
2 COMPUTING, 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 assert!(
114 !self.running.get(),
115 "reactive cycle: a memo depends on itself"
116 );
117 if !self.stale.get() {
118 return;
119 }
120 self.running.set(true);
121 COMPUTING.with(|depth| depth.set(depth.get() + 1));
122 let _evaluation = EvaluationGuard(&self.running);
123 let graph = self.graph();
124 if !self.needs_compute.get() && !untrack(|| graph.observer.changed()) {
125 self.stale.set(false);
126 return;
127 }
128 self.needs_compute.set(true);
131 let next = {
132 let _run = graph.observer.begin();
133 let _tracking = TrackingGuard::replace(Some(Rc::downgrade(&graph.observer)));
134 (self.compute)()
135 };
136 let equal = untrack(|| {
137 self.value
138 .borrow()
139 .as_ref()
140 .is_some_and(|old| (self.equal)(old, &next))
141 });
142 if equal {
143 untrack(|| drop(next));
144 } else {
145 let old = self.value.replace(Some(next));
146 graph.source.advance();
147 untrack(|| drop(old));
149 }
150 self.needs_compute.set(false);
151 self.stale.set(false);
152 }
153
154 fn invalidate(&self) -> &Source {
155 self.stale.set(true);
156 &self.graph.get().expect("memo graph initialized").source
158 }
159}
160
161impl<T: 'static> Memo<T> {
162 pub fn with<R>(&self, read: impl FnOnce(&T) -> R) -> R {
165 self.0.refresh();
166 track(&self.0.graph.get().expect("memo graph initialized").source);
168 read(self.0.value.borrow().as_ref().expect("memo evaluated"))
169 }
170
171 pub fn with_untracked<R>(&self, read: impl FnOnce(&T) -> R) -> R {
173 untrack(|| self.with(read))
174 }
175}
176
177impl<T: Clone + 'static> Memo<T> {
178 pub fn get(&self) -> T {
179 self.with(Clone::clone)
180 }
181 pub fn get_untracked(&self) -> T {
182 self.with_untracked(Clone::clone)
183 }
184}