use std::{
panic::{catch_unwind, AssertUnwindSafe},
ptr::NonNull,
sync::{
atomic::{AtomicBool, Ordering},
Arc, Mutex,
},
};
use moirai_core::{
error::{ExecutorError, ExecutorResult, TaskError},
executor::ExecutorConfig,
task::{TaskHandle, TaskId, TaskResultSender},
Priority,
};
use crate::{
metrics::ExecutorMetrics,
registry::{TaskLifecycleToken, TaskRegistry},
schedule::{SchedulerScope, SyncTask, ThreadScheduler, WorkClass, WorkScheduler},
};
mod async_state;
pub(crate) mod control;
pub(crate) mod manager;
pub(crate) mod spawner;
#[cfg(test)]
mod tests;
#[derive(Clone, Copy)]
struct MetricsRef {
metrics: NonNull<ExecutorMetrics>,
}
unsafe impl Send for MetricsRef {}
impl MetricsRef {
#[inline]
fn new(metrics: &Arc<ExecutorMetrics>) -> Self {
Self {
metrics: NonNull::from(metrics.as_ref()),
}
}
#[inline]
fn get(self) -> &'static ExecutorMetrics {
unsafe { self.metrics.as_ref() }
}
}
pub struct HybridExecutor<S: WorkScheduler = ThreadScheduler> {
config: ExecutorConfig,
scheduler: S,
task_registry: Arc<Mutex<TaskRegistry>>,
metrics: Arc<ExecutorMetrics>,
shutdown_signal: Arc<AtomicBool>,
}
impl HybridExecutor<ThreadScheduler> {
pub fn new(config: ExecutorConfig) -> ExecutorResult<Self> {
let scheduler = ThreadScheduler::new(config.worker_threads, &config.thread_name_prefix)?;
let metrics = Arc::new(ExecutorMetrics::new());
metrics.update_worker_counts(0, scheduler.worker_count(), scheduler.worker_count());
Ok(Self {
config,
scheduler,
task_registry: Arc::new(Mutex::new(TaskRegistry::new())),
metrics,
shutdown_signal: Arc::new(AtomicBool::new(false)),
})
}
pub fn scope<'scope, C, F>(&'scope self, body: F) -> ExecutorResult<()>
where
C: WorkClass,
F: FnOnce(&SchedulerScope<'scope, C>) -> ExecutorResult<()>,
{
self.scheduler.scope::<C, _>(Priority::Normal, None, body)
}
}
impl<S: WorkScheduler> HybridExecutor<S> {
pub fn config(&self) -> &ExecutorConfig {
&self.config
}
pub fn shutdown(&mut self) -> ExecutorResult<()> {
self.shutdown_signal.store(true, Ordering::Release);
self.scheduler.shutdown();
Ok(())
}
pub fn metrics(&self) -> &ExecutorMetrics {
self.refresh_scheduler_metrics();
&self.metrics
}
pub fn submit_task<F>(&self, task: F) -> ExecutorResult<TaskId>
where
F: FnOnce() + Send + 'static,
{
self.spawn_result::<SyncTask, _>(Priority::Normal, None, task)
.map(|handle| handle.id())
}
pub(super) fn spawn_result<C, R>(
&self,
priority: Priority,
locality_hint: Option<usize>,
func: impl FnOnce() -> R + Send + 'static,
) -> ExecutorResult<TaskHandle<R>>
where
C: WorkClass,
R: Send + 'static,
{
let (task_id, lifecycle) = self.register_task(priority)?;
let (handle, result_sender) = TaskHandle::new_pending(task_id);
let metrics = MetricsRef::new(&self.metrics);
self.scheduler
.schedule::<C, _>(priority, locality_hint, move |worker_id| {
let Some(running) = lifecycle.start_unless_cancelled(worker_id) else {
metrics.get().record_task_cancelled();
result_sender.send(Err(TaskError::Cancelled));
return;
};
let result = catch_unwind(AssertUnwindSafe(func));
let execution_time = running.complete();
send_task_result(result, result_sender, metrics.get(), execution_time);
})?;
self.metrics.record_task_spawned();
Ok(handle)
}
pub fn for_each_indexed<'scope, C, F>(&'scope self, count: usize, task: F) -> ExecutorResult<()>
where
C: WorkClass,
F: Fn(usize) + Send + Sync + 'scope,
{
self.scheduler
.for_each_indexed::<C, _>(Priority::Normal, None, count, task)
}
pub fn map_reduce_indexed<'scope, C, T, Map, Reduce>(
&'scope self,
count: usize,
identity: T,
map: Map,
reduce: Reduce,
) -> ExecutorResult<T>
where
C: WorkClass,
T: Send + Clone + 'scope,
Map: Fn(usize) -> T + Send + Sync + 'scope,
Reduce: Fn(T, T) -> T + Send + Sync + 'scope,
{
self.scheduler.map_reduce_indexed::<C, _, _, _>(
Priority::Normal,
None,
count,
identity,
map,
reduce,
)
}
pub fn active_workers(&self) -> usize {
self.scheduler.active_workers()
}
pub fn total_workers(&self) -> usize {
self.scheduler.worker_count()
}
pub fn pending_tasks(&self) -> usize {
self.scheduler.pending_tasks()
}
pub fn has_work(&self) -> bool {
self.scheduler.has_work()
}
pub fn join(&self) -> ExecutorResult<()> {
self.scheduler.join()?;
self.refresh_scheduler_metrics();
Ok(())
}
fn register_task(&self, priority: Priority) -> ExecutorResult<(TaskId, TaskLifecycleToken)> {
let mut registry = self.task_registry.lock().map_err(|_| {
ExecutorError::ResourceExhausted("task registry lock poisoned".to_string())
})?;
let (task_id, lifecycle) = registry.register_next_task();
lifecycle.set_priority(priority);
Ok((TaskId::new(task_id), lifecycle))
}
fn refresh_scheduler_metrics(&self) {
let snapshot = self.scheduler.metrics();
self.metrics.update_worker_counts(
snapshot.active_workers,
snapshot
.worker_count
.saturating_sub(snapshot.active_workers),
snapshot.worker_count,
);
self.metrics.update_queue_metrics(snapshot.pending_tasks);
}
}
impl<S: WorkScheduler> Drop for HybridExecutor<S> {
fn drop(&mut self) {
let _ = Self::shutdown(self);
}
}
#[inline]
fn send_task_result<T>(
result: Result<T, Box<dyn std::any::Any + Send>>,
sender: TaskResultSender<T>,
metrics: &ExecutorMetrics,
execution_time: core::time::Duration,
) where
T: Send + 'static,
{
match result {
Ok(value) => {
sender.send(Ok(value));
metrics.record_task_completed(execution_time);
}
Err(_) => {
sender.send(Err(TaskError::Panicked));
metrics.record_task_failed();
}
}
}