Skip to main content

wasmtime/runtime/component/concurrent/
task_group_hook.rs

1use crate::component::concurrent::table::{TableDebug, TableId};
2use crate::component::concurrent::{ConcurrentState, CurrentThread};
3use crate::error::Result;
4use crate::store::StoreOpaque;
5use crate::{AsContextMut as _, Store, StoreContextMut};
6use alloc::boxed::Box;
7use alloc::vec::Vec;
8
9struct TaskGroup {
10    ref_count: usize,
11}
12
13impl TableDebug for TaskGroup {
14    fn type_name() -> &'static str {
15        "TaskGroup"
16    }
17}
18
19/// Represents a "task group" containing the "root" task of a host->guest call,
20/// plus any subtasks transitively created by that task.
21///
22/// See [TaskGroupHook] for details.
23#[derive(Copy, Clone, PartialEq, Eq, Hash, PartialOrd, Ord, Debug)]
24pub struct TaskGroupId(TableId<TaskGroup>);
25
26/// Trait for being notified by the runtime of activity concerning a "task group".
27///
28/// The Component Model specification has no notion of a "task group"[^1], but
29/// we define one here in order to enable embedders to associate guest->host
30/// calls with corresponding host->guest calls in a predictable way.
31///
32/// A `TaskGroupId` is allocated whenever a guest task is created for a
33/// host->guest call, and `handle_start` is called.  That task is considered the
34/// "root" task for the task group, and any subtasks transitively created by it
35/// will also be considered part of that task group.  Whenever the runtime
36/// switches (Component-Model-level) threads, it will call `handle_exit` for the
37/// group to which the old thread belonged, if any, and call `handle_enter` for
38/// the group to which the new thread belongs.  Only once all the threads of all
39/// those tasks have exited (and the guest has dropped any subtask handles
40/// referring to any of those tasks) will the `TaskGroupId` be deallocated, at
41/// which point `handle_finish` will be called.
42///
43/// Note that a given `TaskGroupId` may be reused after `handle_finish` is
44/// called, so implementations of this trait must take care to reset any state
45/// associated with it.
46///
47/// Each of these functions may return an error, in which case any running guest
48/// code will trap, the store will be poisoned such that it cannot be used to
49/// run any further guest code, and the error will propagate back to the
50/// host->guest caller.  Furthermore, if and when the store is poisoned due to
51/// any unrecoverable error (whether it was produced by the hook, the guest, or
52/// the host), the currently-entered group will be exited (i.e. `handle_exit`
53/// called), if any, and any started group will be finished
54/// (i.e. `handle_finish` called).
55///
56/// As of this writing,
57/// <https://github.com/WebAssembly/component-model/pull/730> (which adds
58/// `thread.set-task` and related intrinsics) has not yet been merged.  Once it
59/// has, and Wasmtime adds support for that feature, it will be possible for
60/// guest threads to change their task; in that case, the thread will
61/// effectively join whatever group the new task belongs to, which might not be
62/// the same as that of the old task.  In addition, the new `thread.get-task`
63/// intrinsic will give the guest another way (besides subtask handles) to keep
64/// tasks alive beyond the point when all their threads have exited or switched
65/// tasks, in which case the group it belongs to will not be disposed until all
66/// such tasks have been dropped using `task.drop`.
67///
68/// [^1]: Although it does _imply_ such a notion in the discussion of ["semantic
69/// tail
70/// calls"](https://github.com/WebAssembly/component-model/blob/d1daf829e2da2293091c105121383ffbc3b3515b/design/mvp/Concurrency.md?plain=1#L339-L3410).
71pub trait TaskGroupHook: Send + Sync + 'static {
72    /// Handle notification that a new task group has been created (i.e. a
73    /// host->guest call has been prepared).
74    fn handle_start(&mut self, id: TaskGroupId) -> Result<()>;
75    /// Handle notification that the runtime has switched to a thread belonging to
76    /// the specified task group.
77    fn handle_enter(&mut self, id: TaskGroupId) -> Result<()>;
78    /// Handle notification that the runtime has switched away from a thread
79    /// belonging to the specified task group.
80    fn handle_exit(&mut self, id: TaskGroupId) -> Result<()>;
81    /// Handle notification that the specified task group has been disposed of
82    /// (i.e. the task created for the host->guest call for which the task group
83    /// was created has exited, along with any and all subtasks transitively
84    /// created by that task, and the guest has dropped any and all handles to
85    /// those tasks).
86    fn handle_finish(&mut self, id: TaskGroupId) -> Result<()>;
87}
88
89impl<T> Store<T> {
90    /// Convenience wrapper for [`StoreContextMut::task_group_hook`]
91    pub fn task_group_hook(&mut self, hook: impl TaskGroupHook) {
92        self.as_context_mut().task_group_hook(hook);
93    }
94}
95
96impl<T> StoreContextMut<'_, T> {
97    /// Set a [`TaskGroupHook`] for this store.
98    ///
99    /// This will overwrite any hook that was previously set.
100    pub fn task_group_hook(self, hook: impl TaskGroupHook) {
101        self.0
102            .concurrent_state_mut_without_forcing_current_thread()
103            .task_group_hook = Some(Box::new(hook));
104    }
105}
106
107impl StoreOpaque {
108    pub(crate) fn clean_up_task_groups(&mut self) {
109        if !self.concurrency_support() {
110            return;
111        }
112
113        // Note that we ignore all errors here since, if we've reached here,
114        // it's either because we've already poisoned the store and are in the
115        // process of propagating the error which caused it to be poisoned or
116        // because we're disposing of the store entirely.
117
118        let state = self.concurrent_state_mut_without_forcing_current_thread();
119        if let Some(mut hook) = state.task_group_hook.take() {
120            let thread = state.unforced_current_thread;
121            if let Ok(Some(group)) = thread.group(state) {
122                _ = hook.handle_exit(group);
123            }
124
125            let groups = state
126                .table
127                .get_mut()
128                .iter_mut()
129                .filter_map(|(rep, entry)| {
130                    if entry.downcast_mut::<TaskGroup>().is_some() {
131                        Some(TableId::new(rep))
132                    } else {
133                        None
134                    }
135                })
136                .collect::<Vec<_>>();
137
138            for group in groups {
139                state.delete(group).unwrap();
140                _ = hook.handle_finish(TaskGroupId(group));
141            }
142
143            state.task_group_hook = Some(hook);
144        }
145    }
146}
147
148impl CurrentThread {
149    fn group(&self, state: &mut ConcurrentState) -> Result<Option<TaskGroupId>> {
150        Ok(match self {
151            Self::Guest(thread) | Self::DeferredHost(thread) => {
152                Some(state.get_mut(thread.task)?.group)
153            }
154            Self::Host(task) => Some(state.get_mut(*task)?.group),
155            Self::None => None,
156        })
157    }
158}
159
160impl ConcurrentState {
161    pub(super) fn handle_thread_switch(
162        &mut self,
163        old: CurrentThread,
164        new: CurrentThread,
165    ) -> Result<()> {
166        let old_group = old.group(self)?;
167        let new_group = new.group(self)?;
168        if let (true, Some(hook)) = ((old_group != new_group), &mut self.task_group_hook) {
169            if let Some(group) = old_group {
170                hook.handle_exit(group)?;
171            }
172
173            if let Some(group) = new_group {
174                hook.handle_enter(group)?;
175            }
176        }
177        Ok(())
178    }
179
180    pub(super) fn make_task_group(&mut self) -> Result<TaskGroupId> {
181        let group = TaskGroupId(self.push(TaskGroup { ref_count: 1 })?);
182        if let Some(hook) = &mut self.task_group_hook {
183            hook.handle_start(group)?;
184        }
185        log::trace!("new {group:?}");
186        Ok(group)
187    }
188
189    pub(super) fn increment_group_ref_count(&mut self, group: TaskGroupId) -> Result<()> {
190        let count = &mut self.get_mut(group.0)?.ref_count;
191        *count += 1;
192        log::trace!("increment {group:?} to {count}");
193        Ok(())
194    }
195
196    pub(super) fn decrement_group_ref_count(&mut self, group: TaskGroupId) -> Result<()> {
197        let count = &mut self.get_mut(group.0)?.ref_count;
198        assert!(*count >= 1);
199        *count -= 1;
200        log::trace!("decrement {group:?} to {count}");
201        if *count == 0 {
202            self.delete(group.0)?;
203            if let Some(hook) = &mut self.task_group_hook {
204                hook.handle_finish(group)?;
205            }
206        }
207        Ok(())
208    }
209}