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}