1use std::{
10 panic::{AssertUnwindSafe, catch_unwind},
11 ptr::NonNull,
12 sync::{
13 Arc,
14 atomic::{AtomicBool, Ordering},
15 },
16};
17
18use moirai_core::{
19 Priority,
20 error::{ExecutorResult, TaskError},
21 executor::ExecutorConfig,
22 task::{TaskHandle, TaskId, TaskResultSender},
23};
24
25use crate::{
26 metrics::ExecutorMetrics,
27 registry::{RetentionPolicy, SchedulerStateLease, TaskLifecycleToken, TaskRegistry},
28 schedule::{SchedulerScope, SyncTask, ThreadScheduler, WorkClass, WorkScheduler},
29};
30
31mod async_state;
32pub(crate) mod control;
33pub(crate) mod manager;
34#[cfg(test)]
35mod retention_tests;
36pub(crate) mod spawner;
37#[cfg(test)]
38mod tests;
39
40#[derive(Clone, Copy)]
41struct MetricsRef {
42 metrics: NonNull<ExecutorMetrics>,
43}
44
45unsafe impl Send for MetricsRef {}
49
50impl MetricsRef {
51 #[inline]
52 fn new(metrics: &Arc<ExecutorMetrics>) -> Self {
53 Self {
54 metrics: NonNull::from(metrics.as_ref()),
55 }
56 }
57
58 #[inline]
59 fn get(self) -> &'static ExecutorMetrics {
60 unsafe { self.metrics.as_ref() }
63 }
64}
65
66pub struct HybridExecutor<S: WorkScheduler = ThreadScheduler> {
74 config: ExecutorConfig,
75 scheduler: S,
76 task_registry: Arc<TaskRegistry>,
77 metrics: Arc<ExecutorMetrics>,
78 shutdown_signal: Arc<AtomicBool>,
79}
80
81impl HybridExecutor<ThreadScheduler> {
82 pub fn new(config: ExecutorConfig) -> ExecutorResult<Self> {
94 let scheduler = ThreadScheduler::from_executor_config(&config)?;
95 let task_registry = Arc::new(
96 RetentionPolicy::from_cleanup(&config.cleanup)
97 .map_or_else(TaskRegistry::new, TaskRegistry::with_retention),
98 );
99 let metrics = Arc::new(ExecutorMetrics::new());
100 scheduler.retain_lifetime_owner((Arc::clone(&task_registry), Arc::clone(&metrics)));
101 metrics.update_worker_counts(0, scheduler.worker_count(), scheduler.worker_count());
102
103 Ok(Self {
104 config,
105 scheduler,
106 task_registry,
107 metrics,
108 shutdown_signal: Arc::new(AtomicBool::new(false)),
109 })
110 }
111
112 pub fn scope<'scope, C, F>(&'scope self, body: F) -> ExecutorResult<()>
123 where
124 C: WorkClass,
125 F: FnOnce(&SchedulerScope<'scope, C>) -> ExecutorResult<()>,
126 {
127 self.scheduler.scope::<C, _>(Priority::Normal, None, body)
128 }
129}
130
131impl<S: WorkScheduler> HybridExecutor<S> {
132 pub fn config(&self) -> &ExecutorConfig {
134 &self.config
135 }
136
137 pub fn shutdown(&mut self) -> ExecutorResult<()> {
139 self.shutdown_signal.store(true, Ordering::Release);
140 self.scheduler.shutdown();
141 Ok(())
142 }
143
144 pub fn metrics(&self) -> &ExecutorMetrics {
146 self.refresh_scheduler_metrics();
147 &self.metrics
148 }
149
150 pub fn submit_task<F>(&self, task: F) -> ExecutorResult<TaskId>
152 where
153 F: FnOnce() + Send + 'static,
154 {
155 self.spawn_result::<SyncTask, _>(Priority::Normal, None, task)
156 .map(|handle| handle.id())
157 }
158
159 pub(super) fn spawn_result<C, R>(
167 &self,
168 priority: Priority,
169 locality_hint: Option<usize>,
170 func: impl FnOnce() -> R + Send + 'static,
171 ) -> ExecutorResult<TaskHandle<R>>
172 where
173 C: WorkClass,
174 R: Send + 'static,
175 {
176 let (task_id, lifecycle) = self.register_scheduled_task(priority)?;
177
178 let (handle, result_sender) = TaskHandle::new_pending(task_id);
179 let metrics = MetricsRef::new(&self.metrics);
180
181 self.scheduler
182 .schedule::<C, _>(priority, locality_hint, move |worker_id| {
183 let Some(running) = lifecycle.start_unless_cancelled(worker_id) else {
184 metrics.get().record_task_cancelled();
187 result_sender.send(Err(TaskError::Cancelled));
188 return;
189 };
190 let result = catch_unwind(AssertUnwindSafe(func));
191 let execution_time = running.complete();
192 send_task_result(result, result_sender, metrics.get(), execution_time);
193 })?;
194
195 self.metrics.record_task_spawned();
196 Ok(handle)
197 }
198
199 pub fn for_each_indexed<'scope, C, F>(&'scope self, count: usize, task: F) -> ExecutorResult<()>
204 where
205 C: WorkClass,
206 F: Fn(usize) + Send + Sync + 'scope,
207 {
208 self.scheduler
209 .for_each_indexed::<C, _>(Priority::Normal, None, count, task)
210 }
211
212 pub fn map_reduce_indexed<'scope, C, T, Map, Reduce>(
214 &'scope self,
215 count: usize,
216 identity: T,
217 map: Map,
218 reduce: Reduce,
219 ) -> ExecutorResult<T>
220 where
221 C: WorkClass,
222 T: Send + Clone + 'scope,
223 Map: Fn(usize) -> T + Send + Sync + 'scope,
224 Reduce: Fn(T, T) -> T + Send + Sync + 'scope,
225 {
226 self.scheduler.map_reduce_indexed::<C, _, _, _>(
227 Priority::Normal,
228 None,
229 count,
230 identity,
231 map,
232 reduce,
233 )
234 }
235
236 pub fn active_workers(&self) -> usize {
238 self.scheduler.active_workers()
239 }
240
241 pub fn total_workers(&self) -> usize {
243 self.scheduler.worker_count()
244 }
245
246 pub fn pending_tasks(&self) -> usize {
248 self.scheduler.pending_tasks()
249 }
250
251 pub fn has_work(&self) -> bool {
253 self.scheduler.has_work()
254 }
255
256 pub fn join(&self) -> ExecutorResult<()> {
258 self.scheduler.join()?;
259 self.refresh_scheduler_metrics();
260 Ok(())
261 }
262
263 fn register_scheduled_task(
264 &self,
265 priority: Priority,
266 ) -> ExecutorResult<(TaskId, TaskLifecycleToken<SchedulerStateLease>)> {
267 let registry = &self.task_registry;
268 let (task_id, lifecycle) = unsafe { registry.register_next_scheduled_task() };
282 lifecycle.set_priority(priority);
283 Ok((TaskId::new(task_id), lifecycle))
284 }
285
286 fn refresh_scheduler_metrics(&self) {
287 let snapshot = self.scheduler.metrics();
288 self.metrics.update_worker_counts(
289 snapshot.active_workers,
290 snapshot
291 .worker_count
292 .saturating_sub(snapshot.active_workers),
293 snapshot.worker_count,
294 );
295 self.metrics.update_queue_metrics(snapshot.pending_tasks);
296 }
297}
298
299impl<S: WorkScheduler> Drop for HybridExecutor<S> {
300 fn drop(&mut self) {
301 let _ = Self::shutdown(self);
302 }
303}
304
305#[inline]
306fn send_task_result<T>(
307 result: Result<T, Box<dyn std::any::Any + Send>>,
308 sender: TaskResultSender<T>,
309 metrics: &ExecutorMetrics,
310 execution_time: core::time::Duration,
311) where
312 T: Send + 'static,
313{
314 match result {
315 Ok(value) => {
316 sender.send(Ok(value));
317 metrics.record_task_completed(execution_time);
318 }
319 Err(_) => {
320 sender.send(Err(TaskError::Panicked));
321 metrics.record_task_failed();
322 }
323 }
324}