1use 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
43unsafe 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 unsafe { self.metrics.as_ref() }
62 }
63}
64
65pub 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 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 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 pub fn config(&self) -> &ExecutorConfig {
118 &self.config
119 }
120
121 pub fn shutdown(&mut self) -> ExecutorResult<()> {
123 self.shutdown_signal.store(true, Ordering::Release);
124 self.scheduler.shutdown();
125 Ok(())
126 }
127
128 pub fn metrics(&self) -> &ExecutorMetrics {
130 self.refresh_scheduler_metrics();
131 &self.metrics
132 }
133
134 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 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 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 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 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 pub fn active_workers(&self) -> usize {
222 self.scheduler.active_workers()
223 }
224
225 pub fn total_workers(&self) -> usize {
227 self.scheduler.worker_count()
228 }
229
230 pub fn pending_tasks(&self) -> usize {
232 self.scheduler.pending_tasks()
233 }
234
235 pub fn has_work(&self) -> bool {
237 self.scheduler.has_work()
238 }
239
240 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}