Skip to main content

lgui_core/core/component/
context_value.rs

1use std::{
2    any::{type_name, Any, TypeId},
3    cell::{Cell, RefCell},
4    collections::{HashMap, HashSet},
5    rc::Rc,
6};
7
8use super::{
9    ComponentId, ComponentTree, EffectRegistry, HookId, HookSlotKind, IntoEffectCleanup, UiEffect,
10};
11
12#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
13struct ProviderKey {
14    component: ComponentId,
15    value_type: TypeId,
16}
17
18thread_local! {
19    static CURRENT_CONTEXT: RefCell<Vec<CurrentContext>> = const { RefCell::new(Vec::new()) };
20}
21
22#[derive(Clone)]
23struct CurrentContext {
24    registry: ContextRegistry,
25    consumer: ComponentId,
26}
27
28struct StagedListenerEffect {
29    component: ComponentId,
30    index: usize,
31    deps: Box<dyn Any>,
32    deps_equal: fn(&dyn Any, &dyn Any) -> bool,
33    run: Box<dyn FnOnce() -> Option<UiEffect> + 'static>,
34}
35
36#[derive(Clone, Default)]
37pub struct ContextRegistry {
38    inner: Rc<ContextRegistryState>,
39}
40
41#[derive(Default)]
42struct ContextRegistryState {
43    providers: RefCell<HashMap<ProviderKey, Box<dyn Any>>>,
44    active: RefCell<HashMap<TypeId, Vec<ProviderKey>>>,
45    consumers: RefCell<HashMap<ProviderKey, HashSet<ComponentId>>>,
46    provider_rollback: RefCell<HashMap<ProviderKey, Option<Box<dyn Any>>>>,
47    consumer_rollback: RefCell<Option<HashMap<ProviderKey, HashSet<ComponentId>>>>,
48    listener_counts: RefCell<HashMap<ComponentId, usize>>,
49    listener_count_rollback: RefCell<Option<HashMap<ComponentId, usize>>>,
50    listener_rendered: RefCell<HashSet<ComponentId>>,
51    listener_next: RefCell<HashMap<ComponentId, usize>>,
52    staged_listener_effects: RefCell<Vec<StagedListenerEffect>>,
53    render_active: Cell<bool>,
54}
55
56pub struct ContextProviderGuard<'a> {
57    registry: &'a ContextRegistry,
58    value_type: TypeId,
59}
60
61pub(crate) struct CurrentContextGuard;
62
63/// Reads a typed value from the nearest provider and subscribes the component
64/// currently being rendered to provider changes.
65pub fn use_context<T>() -> T
66where
67    T: Clone + 'static,
68{
69    try_use_context::<T>().unwrap_or_else(|| {
70        panic!(
71            "missing active context value `{}`; context hooks may only run while rendering a component",
72            type_name::<T>()
73        )
74    })
75}
76
77/// Tries to read a typed value from the nearest provider for the component
78/// currently being rendered.
79pub fn try_use_context<T>() -> Option<T>
80where
81    T: Clone + 'static,
82{
83    let current = CURRENT_CONTEXT.with(|stack| stack.borrow().last().cloned())?;
84    current.registry.read(current.consumer)
85}
86
87impl ContextRegistry {
88    pub fn new() -> Self {
89        Self::default()
90    }
91
92    pub fn begin_render(&self) {
93        self.restore_render_state();
94        self.inner.active.borrow_mut().clear();
95        *self.inner.consumer_rollback.borrow_mut() = Some(self.inner.consumers.borrow().clone());
96        *self.inner.listener_count_rollback.borrow_mut() =
97            Some(self.inner.listener_counts.borrow().clone());
98        self.inner.listener_rendered.borrow_mut().clear();
99        self.inner.listener_next.borrow_mut().clear();
100        self.inner.staged_listener_effects.borrow_mut().clear();
101        self.inner.render_active.set(true);
102    }
103
104    pub fn begin_component(&self, component: ComponentId) {
105        self.inner.consumers.borrow_mut().retain(|_, consumers| {
106            consumers.remove(&component);
107            !consumers.is_empty()
108        });
109        self.inner.listener_rendered.borrow_mut().insert(component);
110        self.inner.listener_next.borrow_mut().insert(component, 0);
111    }
112
113    pub(crate) fn validate_listener_hooks(&self) {
114        let rendered = self.inner.listener_rendered.borrow();
115        let next = self.inner.listener_next.borrow();
116        let counts = self.inner.listener_counts.borrow();
117        for component in rendered.iter().copied() {
118            let current = next.get(&component).copied().unwrap_or_default();
119            if let Some(previous) = counts.get(&component) {
120                assert_eq!(
121                    *previous, current,
122                    "component {component} changed its receiver-free listener hook count from {previous} to {current}"
123                );
124            }
125        }
126        drop(counts);
127        drop(next);
128        drop(rendered);
129    }
130
131    pub(crate) fn commit_listener_effects(
132        &self,
133        components: &ComponentTree,
134        effects: &EffectRegistry,
135    ) {
136        for staged in self.inner.staged_listener_effects.borrow_mut().drain(..) {
137            if !components.is_alive(staged.component) {
138                continue;
139            }
140            effects.register_erased(
141                HookId::new(staged.component, staged.index, HookSlotKind::Listener),
142                staged.deps,
143                staged.deps_equal,
144                staged.run,
145            );
146        }
147    }
148
149    pub fn end_render(&self, components: &ComponentTree) {
150        debug_assert!(
151            self.inner.active.borrow().values().all(Vec::is_empty),
152            "context provider stack was not balanced"
153        );
154        self.inner.active.borrow_mut().clear();
155        self.inner
156            .providers
157            .borrow_mut()
158            .retain(|key, _| components.is_alive(key.component));
159        self.inner.consumers.borrow_mut().retain(|key, consumers| {
160            if !components.is_alive(key.component) {
161                return false;
162            }
163            consumers.retain(|consumer| components.is_alive(*consumer));
164            !consumers.is_empty()
165        });
166        self.inner.provider_rollback.borrow_mut().clear();
167        self.inner.consumer_rollback.borrow_mut().take();
168        {
169            let rendered = self.inner.listener_rendered.borrow();
170            let next = self.inner.listener_next.borrow();
171            let mut counts = self.inner.listener_counts.borrow_mut();
172            for component in rendered.iter().copied() {
173                counts.insert(component, next.get(&component).copied().unwrap_or_default());
174            }
175            counts.retain(|component, _| components.is_alive(*component));
176        }
177        self.inner.listener_count_rollback.borrow_mut().take();
178        self.inner.listener_rendered.borrow_mut().clear();
179        self.inner.listener_next.borrow_mut().clear();
180        self.inner.staged_listener_effects.borrow_mut().clear();
181        self.inner.render_active.set(false);
182    }
183
184    pub fn abort_render(&self, _components: &ComponentTree) {
185        self.inner.active.borrow_mut().clear();
186        self.inner.staged_listener_effects.borrow_mut().clear();
187        self.inner.listener_rendered.borrow_mut().clear();
188        self.inner.listener_next.borrow_mut().clear();
189        if let Some(counts) = self.inner.listener_count_rollback.borrow_mut().take() {
190            *self.inner.listener_counts.borrow_mut() = counts;
191        }
192        self.restore_render_state();
193        self.inner.render_active.set(false);
194    }
195
196    pub fn provide<T>(
197        &self,
198        owner: ComponentId,
199        value: T,
200        components: &ComponentTree,
201    ) -> ContextProviderGuard<'_>
202    where
203        T: Clone + PartialEq + 'static,
204    {
205        let value_type = TypeId::of::<T>();
206        let key = ProviderKey {
207            component: owner,
208            value_type,
209        };
210        let changed = {
211            let mut providers = self.inner.providers.borrow_mut();
212            if self.inner.render_active.get()
213                && !self.inner.provider_rollback.borrow().contains_key(&key)
214            {
215                let previous = providers.get(&key).map(|current| {
216                    Box::new(
217                        current
218                            .downcast_ref::<T>()
219                            .unwrap_or_else(|| {
220                                panic!("context provider type mismatch for `{}`", type_name::<T>())
221                            })
222                            .clone(),
223                    ) as Box<dyn Any>
224                });
225                self.inner
226                    .provider_rollback
227                    .borrow_mut()
228                    .insert(key, previous);
229            }
230            match providers.get_mut(&key) {
231                Some(current) => {
232                    let current = current.downcast_mut::<T>().unwrap_or_else(|| {
233                        panic!("context provider type mismatch for `{}`", type_name::<T>())
234                    });
235                    if current == &value {
236                        false
237                    } else {
238                        *current = value;
239                        true
240                    }
241                }
242                None => {
243                    providers.insert(key, Box::new(value));
244                    true
245                }
246            }
247        };
248        if changed {
249            if let Some(consumers) = self.inner.consumers.borrow().get(&key) {
250                for consumer in consumers {
251                    components.mark_dirty(*consumer);
252                }
253            }
254        }
255        self.inner
256            .active
257            .borrow_mut()
258            .entry(value_type)
259            .or_default()
260            .push(key);
261        ContextProviderGuard {
262            registry: self,
263            value_type,
264        }
265    }
266
267    pub fn read<T>(&self, consumer: ComponentId) -> Option<T>
268    where
269        T: Clone + 'static,
270    {
271        let value_type = TypeId::of::<T>();
272        let key = self
273            .inner
274            .active
275            .borrow()
276            .get(&value_type)
277            .and_then(|providers| providers.last())
278            .copied()?;
279        self.inner
280            .consumers
281            .borrow_mut()
282            .entry(key)
283            .or_default()
284            .insert(consumer);
285        Some(
286            self.inner
287                .providers
288                .borrow()
289                .get(&key)
290                .and_then(|value| value.downcast_ref::<T>())
291                .unwrap_or_else(|| panic!("context value type mismatch for `{}`", type_name::<T>()))
292                .clone(),
293        )
294    }
295
296    pub fn clear(&self) {
297        self.inner.providers.borrow_mut().clear();
298        self.inner.active.borrow_mut().clear();
299        self.inner.consumers.borrow_mut().clear();
300        self.inner.provider_rollback.borrow_mut().clear();
301        self.inner.consumer_rollback.borrow_mut().take();
302        self.inner.listener_counts.borrow_mut().clear();
303        self.inner.listener_count_rollback.borrow_mut().take();
304        self.inner.listener_rendered.borrow_mut().clear();
305        self.inner.listener_next.borrow_mut().clear();
306        self.inner.staged_listener_effects.borrow_mut().clear();
307        self.inner.render_active.set(false);
308    }
309
310    fn stage_current_listener<D, F, R>(&self, component: ComponentId, deps: D, effect: F)
311    where
312        D: Clone + PartialEq + 'static,
313        F: FnOnce() -> R + 'static,
314        R: IntoEffectCleanup,
315    {
316        assert!(
317            self.inner.render_active.get(),
318            "receiver-free `listen` may only be called while rendering a component"
319        );
320        let index = {
321            let mut next = self.inner.listener_next.borrow_mut();
322            let index = next.get(&component).copied().unwrap_or_default();
323            next.insert(component, index + 1);
324            index
325        };
326        self.inner
327            .staged_listener_effects
328            .borrow_mut()
329            .push(StagedListenerEffect {
330                component,
331                index,
332                deps: Box::new(deps),
333                deps_equal: listener_deps_equal::<D>,
334                run: Box::new(move || effect().into_cleanup()),
335            });
336    }
337
338    pub(crate) fn enter_current(&self, consumer: ComponentId) -> CurrentContextGuard {
339        CURRENT_CONTEXT.with(|stack| {
340            stack.borrow_mut().push(CurrentContext {
341                registry: self.clone(),
342                consumer,
343            });
344        });
345        CurrentContextGuard
346    }
347
348    fn restore_render_state(&self) {
349        let rollback = std::mem::take(&mut *self.inner.provider_rollback.borrow_mut());
350        let mut providers = self.inner.providers.borrow_mut();
351        for (key, previous) in rollback {
352            match previous {
353                Some(previous) => {
354                    providers.insert(key, previous);
355                }
356                None => {
357                    providers.remove(&key);
358                }
359            }
360        }
361        drop(providers);
362        if let Some(consumers) = self.inner.consumer_rollback.borrow_mut().take() {
363            *self.inner.consumers.borrow_mut() = consumers;
364        }
365    }
366}
367
368impl Drop for ContextProviderGuard<'_> {
369    fn drop(&mut self) {
370        let mut active = self.registry.inner.active.borrow_mut();
371        let stack = active
372            .get_mut(&self.value_type)
373            .expect("context provider stack disappeared");
374        stack.pop().expect("context provider stack underflow");
375    }
376}
377
378impl Drop for CurrentContextGuard {
379    fn drop(&mut self) {
380        CURRENT_CONTEXT.with(|stack| {
381            stack
382                .borrow_mut()
383                .pop()
384                .expect("current context stack underflow");
385        });
386    }
387}
388
389pub(crate) fn stage_current_listener<D, F, R>(deps: D, effect: F)
390where
391    D: Clone + PartialEq + 'static,
392    F: FnOnce() -> R + 'static,
393    R: IntoEffectCleanup,
394{
395    let current = CURRENT_CONTEXT
396        .with(|stack| stack.borrow().last().cloned())
397        .unwrap_or_else(|| {
398            panic!("receiver-free `listen` may only be called while rendering a component")
399        });
400    current
401        .registry
402        .stage_current_listener(current.consumer, deps, effect);
403}
404
405fn listener_deps_equal<D>(left: &dyn Any, right: &dyn Any) -> bool
406where
407    D: PartialEq + 'static,
408{
409    let left = left
410        .downcast_ref::<D>()
411        .unwrap_or_else(|| panic!("listener dependency type changed between renders"));
412    let right = right
413        .downcast_ref::<D>()
414        .unwrap_or_else(|| panic!("listener dependency type changed during render"));
415    left == right
416}
417
418#[cfg(test)]
419#[path = "context_value_test.rs"]
420mod tests;