Skip to main content

moirai_executor/hybrid/
mod.rs

1//! Main hybrid executor implementation.
2//!
3//! `HybridExecutor` exposes one public execution surface while delegating sync,
4//! async, and blocking work to one scheduler facade. Sync and async-ready work
5//! use the compute worker pool; blocking work uses the facade's bounded lane.
6//! The work-shape choice is encoded by zero-sized marker types in
7//! `crate::schedule`.
8
9use std::{
10    panic::{catch_unwind, AssertUnwindSafe},
11    ptr::NonNull,
12    sync::{
13        atomic::{AtomicBool, Ordering},
14        Arc, Mutex,
15    },
16};
17
18use moirai_core::{
19    error::{ExecutorError, ExecutorResult, TaskError},
20    executor::ExecutorConfig,
21    task::{TaskHandle, TaskId, TaskResultSender},
22    Priority,
23};
24
25use crate::{
26    metrics::ExecutorMetrics,
27    registry::{TaskLifecycleToken, TaskRegistry},
28    schedule::{SchedulerScope, SyncTask, ThreadScheduler, WorkClass, WorkScheduler},
29};
30
31mod async_state;
32pub(crate) mod control;
33pub(crate) mod manager;
34pub(crate) mod spawner;
35#[cfg(test)]
36mod tests;
37
38#[derive(Clone, Copy)]
39struct MetricsRef {
40    metrics: NonNull<ExecutorMetrics>,
41}
42
43// Safety: `MetricsRef` points at `HybridExecutor.metrics`. The executor owns
44// the scheduler and drains scheduled jobs during shutdown/drop before dropping
45// the metrics allocation, so scheduled synchronous/blocking jobs cannot observe
46// a dangling metrics pointer.
47unsafe impl Send for MetricsRef {}
48
49impl MetricsRef {
50    #[inline]
51    fn new(metrics: &Arc<ExecutorMetrics>) -> Self {
52        Self {
53            metrics: NonNull::from(metrics.as_ref()),
54        }
55    }
56
57    #[inline]
58    fn get(self) -> &'static ExecutorMetrics {
59        // Safety: see the `Send` impl invariant. The returned reference is used
60        // only inside scheduled jobs that complete before executor destruction.
61        unsafe { self.metrics.as_ref() }
62    }
63}
64
65/// Main hybrid executor that coordinates sync, async, and blocking tasks.
66///
67/// Generic over the work-stealing runtime `S` behind the
68/// [`WorkScheduler`] seam; the default
69/// [`ThreadScheduler`] backs the production runtime, while the parameter lets a
70/// substitute (e.g. a single-threaded `wasm32` scheduler) be plugged in without
71/// touching this façade.
72pub struct HybridExecutor<S: WorkScheduler = ThreadScheduler> {
73    config: ExecutorConfig,
74    scheduler: S,
75    task_registry: Arc<Mutex<TaskRegistry>>,
76    metrics: Arc<ExecutorMetrics>,
77    shutdown_signal: Arc<AtomicBool>,
78}
79
80impl HybridExecutor<ThreadScheduler> {
81    /// Create a new hybrid executor with the given configuration.
82    pub fn new(config: ExecutorConfig) -> ExecutorResult<Self> {
83        let scheduler = ThreadScheduler::new(config.worker_threads, &config.thread_name_prefix)?;
84        let metrics = Arc::new(ExecutorMetrics::new());
85        metrics.update_worker_counts(0, scheduler.worker_count(), scheduler.worker_count());
86
87        Ok(Self {
88            config,
89            scheduler,
90            task_registry: Arc::new(Mutex::new(TaskRegistry::new())),
91            metrics,
92            shutdown_signal: Arc::new(AtomicBool::new(false)),
93        })
94    }
95
96    /// Run a scoped fan-out directly on the unified scheduler.
97    ///
98    /// This path is for completion-only work that does not require per-task
99    /// result handles or lifecycle metadata. It preserves borrowing semantics
100    /// by waiting for all spawned jobs before returning.
101    ///
102    /// `scope` is inherent to the default [`ThreadScheduler`] backing because its
103    /// signature exposes a concrete [`SchedulerScope`] borrow handle, which is
104    /// outside the substitutable [`WorkScheduler`]
105    /// seam.
106    pub fn scope<'scope, C, F>(&'scope self, body: F) -> ExecutorResult<()>
107    where
108        C: WorkClass,
109        F: FnOnce(&SchedulerScope<'scope, C>) -> ExecutorResult<()>,
110    {
111        self.scheduler.scope::<C, _>(Priority::Normal, None, body)
112    }
113}
114
115impl<S: WorkScheduler> HybridExecutor<S> {
116    /// Get executor configuration.
117    pub fn config(&self) -> &ExecutorConfig {
118        &self.config
119    }
120
121    /// Shutdown the executor gracefully.
122    pub fn shutdown(&mut self) -> ExecutorResult<()> {
123        self.shutdown_signal.store(true, Ordering::Release);
124        self.scheduler.shutdown();
125        Ok(())
126    }
127
128    /// Get executor metrics.
129    pub fn metrics(&self) -> &ExecutorMetrics {
130        self.refresh_scheduler_metrics();
131        &self.metrics
132    }
133
134    /// Submit an untyped synchronous job.
135    pub fn submit_task<F>(&self, task: F) -> ExecutorResult<TaskId>
136    where
137        F: FnOnce() + Send + 'static,
138    {
139        self.spawn_result::<SyncTask, _>(Priority::Normal, None, task)
140            .map(|handle| handle.id())
141    }
142
143    /// Canonical result-producing spawn path shared by every closure-based
144    /// spawn surface (`spawn_blocking`, `submit_task`, and — via a
145    /// task-executing sibling in `spawner` — the `Task`-typed surfaces).
146    ///
147    /// Registers the task at `priority`, allocates its pending handle, and
148    /// schedules one job that honors queued-task cancellation, contains
149    /// panics, records lifecycle timing, and publishes the result.
150    pub(super) fn spawn_result<C, R>(
151        &self,
152        priority: Priority,
153        locality_hint: Option<usize>,
154        func: impl FnOnce() -> R + Send + 'static,
155    ) -> ExecutorResult<TaskHandle<R>>
156    where
157        C: WorkClass,
158        R: Send + 'static,
159    {
160        let (task_id, lifecycle) = self.register_task(priority)?;
161
162        let (handle, result_sender) = TaskHandle::new_pending(task_id);
163        let metrics = MetricsRef::new(&self.metrics);
164
165        self.scheduler
166            .schedule::<C, _>(priority, locality_hint, move |worker_id| {
167                let Some(running) = lifecycle.start_unless_cancelled(worker_id) else {
168                    // Record before publishing the result so a joiner observes
169                    // the cancelled counter as soon as the handle resolves.
170                    metrics.get().record_task_cancelled();
171                    result_sender.send(Err(TaskError::Cancelled));
172                    return;
173                };
174                let result = catch_unwind(AssertUnwindSafe(func));
175                let execution_time = running.complete();
176                send_task_result(result, result_sender, metrics.get(), execution_time);
177            })?;
178
179        self.metrics.record_task_spawned();
180        Ok(handle)
181    }
182
183    /// Run indexed work in worker-sized chunks on the unified scheduler.
184    ///
185    /// This path avoids per-item task handles and lifecycle metadata when the
186    /// caller only needs completion for a bounded index domain.
187    pub fn for_each_indexed<'scope, C, F>(&'scope self, count: usize, task: F) -> ExecutorResult<()>
188    where
189        C: WorkClass,
190        F: Fn(usize) + Send + Sync + 'scope,
191    {
192        self.scheduler
193            .for_each_indexed::<C, _>(Priority::Normal, None, count, task)
194    }
195
196    /// Run indexed map/reduce in worker-sized chunks on the unified scheduler.
197    pub fn map_reduce_indexed<'scope, C, T, Map, Reduce>(
198        &'scope self,
199        count: usize,
200        identity: T,
201        map: Map,
202        reduce: Reduce,
203    ) -> ExecutorResult<T>
204    where
205        C: WorkClass,
206        T: Send + Clone + 'scope,
207        Map: Fn(usize) -> T + Send + Sync + 'scope,
208        Reduce: Fn(T, T) -> T + Send + Sync + 'scope,
209    {
210        self.scheduler.map_reduce_indexed::<C, _, _, _>(
211            Priority::Normal,
212            None,
213            count,
214            identity,
215            map,
216            reduce,
217        )
218    }
219
220    /// Get the number of active workers.
221    pub fn active_workers(&self) -> usize {
222        self.scheduler.active_workers()
223    }
224
225    /// Get the total number of workers.
226    pub fn total_workers(&self) -> usize {
227        self.scheduler.worker_count()
228    }
229
230    /// Get pending task count across all workers.
231    pub fn pending_tasks(&self) -> usize {
232        self.scheduler.pending_tasks()
233    }
234
235    /// Returns true when queued or active scheduler work exists.
236    pub fn has_work(&self) -> bool {
237        self.scheduler.has_work()
238    }
239
240    /// Wait until queued and active scheduler work completes without shutting down workers.
241    pub fn join(&self) -> ExecutorResult<()> {
242        self.scheduler.join()?;
243        self.refresh_scheduler_metrics();
244        Ok(())
245    }
246
247    fn register_task(&self, priority: Priority) -> ExecutorResult<(TaskId, TaskLifecycleToken)> {
248        let mut registry = self.task_registry.lock().map_err(|_| {
249            ExecutorError::ResourceExhausted("task registry lock poisoned".to_string())
250        })?;
251        let (task_id, lifecycle) = registry.register_next_task();
252        lifecycle.set_priority(priority);
253        Ok((TaskId::new(task_id), lifecycle))
254    }
255
256    fn refresh_scheduler_metrics(&self) {
257        let snapshot = self.scheduler.metrics();
258        self.metrics.update_worker_counts(
259            snapshot.active_workers,
260            snapshot
261                .worker_count
262                .saturating_sub(snapshot.active_workers),
263            snapshot.worker_count,
264        );
265        self.metrics.update_queue_metrics(snapshot.pending_tasks);
266    }
267}
268
269impl<S: WorkScheduler> Drop for HybridExecutor<S> {
270    fn drop(&mut self) {
271        let _ = Self::shutdown(self);
272    }
273}
274
275#[inline]
276fn send_task_result<T>(
277    result: Result<T, Box<dyn std::any::Any + Send>>,
278    sender: TaskResultSender<T>,
279    metrics: &ExecutorMetrics,
280    execution_time: core::time::Duration,
281) where
282    T: Send + 'static,
283{
284    match result {
285        Ok(value) => {
286            sender.send(Ok(value));
287            metrics.record_task_completed(execution_time);
288        }
289        Err(_) => {
290            sender.send(Err(TaskError::Panicked));
291            metrics.record_task_failed();
292        }
293    }
294}