1#![expect(
2 clippy::unwrap_used,
3 reason = "ratchet MOIRAI-UNWRAP-1: pre-existing debt"
4)]
5
6use moirai_core::{Priority, TaskId};
7use moirai_pal::reactor::IoReactor;
8use std::future::Future;
9use std::sync::Arc;
10use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
11use std::task::{Context, Poll, Waker};
12use std::time::Instant;
13
14use crate::executor::handle::AsyncHandle;
15use crate::executor::result_slot::AsyncResultSlot;
16use crate::executor::stats::{AsyncExecutorStats, ExecutorStats};
17use crate::executor::task::{AsyncTask, ErasedTaskFuture};
18
19pub struct AsyncExecutor {
21 reactor: Arc<IoReactor>,
23 run_queue: Arc<moirai_utils::queue::LockFreeQueue<Arc<AsyncTask>>>,
25 stats: AsyncExecutorStats,
27 running: Arc<AtomicBool>,
29 next_task_id: AtomicU64,
31}
32
33impl AsyncExecutor {
34 pub fn new() -> std::io::Result<Self> {
36 let reactor = Arc::new(IoReactor::new()?);
37
38 Ok(Self {
39 reactor,
40 run_queue: Arc::new(moirai_utils::queue::LockFreeQueue::new()),
41 stats: AsyncExecutorStats::default(),
42 running: Arc::new(AtomicBool::new(false)),
43 next_task_id: AtomicU64::new(0),
44 })
45 }
46
47 pub fn spawn<F, T>(&self, future: F) -> AsyncHandle<T>
49 where
50 F: Future<Output = T> + Send + 'static,
51 T: Send + 'static,
52 {
53 self.spawn_with_priority(future, Priority::Normal)
54 }
55
56 pub fn spawn_with_priority<F, T>(&self, future: F, priority: Priority) -> AsyncHandle<T>
58 where
59 F: Future<Output = T> + Send + 'static,
60 T: Send + 'static,
61 {
62 let task_id = TaskId::new(self.next_task_id.fetch_add(1, Ordering::Relaxed));
63 let result_slot = Arc::new(AsyncResultSlot::new());
64 let completion_slot = Arc::clone(&result_slot);
65
66 let wrapped_future = async move {
67 let result = future.await;
68 completion_slot.complete(result);
69 };
70
71 let task = Arc::new(AsyncTask {
72 task_id,
73 future: std::cell::UnsafeCell::new(ErasedTaskFuture::new(wrapped_future)),
74 future_lock: std::sync::Mutex::new(()),
75 run_queue: Arc::downgrade(&self.run_queue),
76 reactor: Arc::downgrade(&self.reactor),
77 is_queued: AtomicBool::new(true),
78 completed: AtomicBool::new(false),
79 priority,
80 created_at: Instant::now(),
81 });
82
83 self.run_queue.enqueue(Arc::clone(&task));
84
85 self.stats.tasks_spawned.fetch_add(1, Ordering::Relaxed);
86 self.stats.tasks_pending.fetch_add(1, Ordering::Relaxed);
87
88 let _ = self.reactor.wake();
89
90 AsyncHandle {
91 task_id,
92 result_slot,
93 }
94 }
95
96 pub fn run(&self) -> std::io::Result<()> {
98 self.running.store(true, Ordering::Relaxed);
105
106 self.reactor.with_active(|| {
107 while self.running.load(Ordering::Acquire) {
114 self.process_pending_tasks();
115
116 let has_tasks = self.stats.tasks_pending.load(Ordering::Acquire) > 0;
117
118 if !has_tasks {
119 if !self.running.load(Ordering::Acquire) {
123 break;
124 }
125 self.reactor.run_iteration(None)?;
126 } else {
127 let run_queue_empty = self.run_queue.is_empty();
128 if run_queue_empty {
129 self.reactor.run_iteration(None)?;
130 } else {
131 self.reactor
132 .run_iteration(Some(std::time::Duration::from_millis(0)))?;
133 }
134 }
135 }
136 Ok(())
137 })
138 }
139
140 pub fn stop(&self) -> std::io::Result<()> {
142 self.running.store(false, Ordering::Release);
149 self.reactor.stop()
150 }
151
152 pub(crate) fn process_pending_tasks(&self) {
154 while let Some(task) = self.run_queue.try_dequeue() {
155 task.is_queued.store(false, Ordering::Relaxed);
161
162 let waker = Waker::from(Arc::clone(&task));
163 let mut context = Context::from_waker(&waker);
164 let task_start = Instant::now();
165
166 let _lock = task.future_lock.lock().unwrap();
182
183 if task.completed.load(Ordering::Acquire) {
184 continue;
185 }
186
187 let future_mut = unsafe { &mut *task.future.get() };
188 match future_mut.poll(&mut context) {
189 std::task::Poll::Ready(()) => {
190 task.completed.store(true, Ordering::Release);
191 self.stats.tasks_completed.fetch_add(1, Ordering::Relaxed);
192 self.stats.tasks_pending.fetch_sub(1, Ordering::Relaxed);
193
194 let execution_time = task_start.elapsed().as_nanos() as u64;
195 self.stats
196 .total_execution_time_ns
197 .fetch_add(execution_time, Ordering::Relaxed);
198 }
199 std::task::Poll::Pending => {}
200 }
201 }
202 }
203
204 pub fn stats(&self) -> ExecutorStats {
206 ExecutorStats {
207 tasks_spawned: self.stats.tasks_spawned.load(Ordering::Relaxed),
208 tasks_completed: self.stats.tasks_completed.load(Ordering::Relaxed),
209 tasks_pending: self.stats.tasks_pending.load(Ordering::Relaxed),
210 total_execution_time_ns: self.stats.total_execution_time_ns.load(Ordering::Relaxed),
211 waker_notifications: self.stats.waker_notifications.load(Ordering::Relaxed),
212 io_operations: self.stats.io_operations.load(Ordering::Relaxed),
213 }
214 }
215
216 pub fn block_on<F, T>(&self, future: F) -> T
218 where
219 F: Future<Output = T> + Send + 'static,
220 T: Send + 'static,
221 {
222 let handle = self.spawn(future);
223 let waker = futures::task::noop_waker();
224 let mut cx = Context::from_waker(&waker);
225 let mut pin_handle = Box::pin(handle);
226
227 self.running.store(true, Ordering::Relaxed);
233
234 loop {
235 self.process_pending_tasks();
236
237 match pin_handle.as_mut().poll(&mut cx) {
238 Poll::Ready(result) => {
239 self.running.store(false, Ordering::Relaxed);
240 return result;
241 }
242 Poll::Pending => {}
243 }
244
245 if self.run_queue.is_empty() {
246 self.reactor.run_iteration(None).ok();
247 } else {
248 self.reactor
249 .run_iteration(Some(std::time::Duration::from_millis(0)))
250 .ok();
251 }
252 }
253 }
254
255 pub fn reactor(&self) -> &IoReactor {
257 &self.reactor
258 }
259}
260
261impl Default for AsyncExecutor {
262 fn default() -> Self {
263 Self::new().expect("Failed to create default AsyncExecutor")
264 }
265}
266
267#[cfg(test)]
268mod tests {
269 use super::*;
270
271 #[test]
272 fn test_native_async_executor_creation() {
273 let executor = AsyncExecutor::new().expect("a fresh AsyncExecutor must build");
274 let stats = executor.stats();
275 assert_eq!(stats.tasks_spawned, 0);
276 assert_eq!(stats.tasks_completed, 0);
277 assert_eq!(stats.tasks_pending, 0);
278 }
279
280 #[test]
281 fn test_task_spawning() {
282 let executor = AsyncExecutor::new().unwrap();
283
284 let _handle = executor.spawn(async { 42 });
285
286 let stats = executor.stats();
287 assert_eq!(stats.tasks_spawned, 1);
288 assert_eq!(stats.tasks_pending, 1);
289 }
290
291 #[test]
292 fn test_task_ids_are_unique() {
293 let executor = AsyncExecutor::new().unwrap();
294
295 let first = executor.spawn(async { 1usize });
296 let second = executor.spawn(async { 2usize });
297
298 assert_ne!(first.id(), second.id());
299 }
300
301 #[test]
302 fn test_ready_task_completion_publishes_result() {
303 use std::task::Poll;
304
305 let executor = AsyncExecutor::new().unwrap();
306 let mut handle = Box::pin(executor.spawn(async { 7usize }));
307 let waker = futures::task::noop_waker();
308 let mut context = Context::from_waker(&waker);
309
310 assert!(matches!(handle.as_mut().poll(&mut context), Poll::Pending));
311
312 executor.process_pending_tasks();
313
314 assert_eq!(executor.stats().tasks_completed, 1);
315 assert!(matches!(handle.as_mut().poll(&mut context), Poll::Ready(7)));
316 }
317
318 #[test]
319 fn test_ready_task_completion_wakes_registered_handle() {
320 use futures::task::{ArcWake, waker_ref};
321 use std::sync::atomic::AtomicUsize;
322 use std::task::Poll;
323
324 struct WakeCounter(AtomicUsize);
325
326 impl ArcWake for WakeCounter {
327 fn wake_by_ref(arc_self: &Arc<Self>) {
328 arc_self.0.fetch_add(1, Ordering::SeqCst);
329 }
330 }
331
332 let executor = AsyncExecutor::new().unwrap();
333 let mut handle = Box::pin(executor.spawn(async { 11usize }));
334 let wake_counter = Arc::new(WakeCounter(AtomicUsize::new(0)));
335 let waker = waker_ref(&wake_counter);
336 let mut context = Context::from_waker(&waker);
337
338 assert!(matches!(handle.as_mut().poll(&mut context), Poll::Pending));
339
340 executor.process_pending_tasks();
341
342 assert_eq!(wake_counter.0.load(Ordering::SeqCst), 1);
343 assert!(matches!(
344 handle.as_mut().poll(&mut context),
345 Poll::Ready(11)
346 ));
347 }
348
349 #[test]
350 fn stale_waker_after_completion_does_not_repoll() {
351 use std::sync::Mutex;
352 use std::sync::atomic::AtomicUsize;
353
354 let executor = AsyncExecutor::new().unwrap();
358 let polls = Arc::new(AtomicUsize::new(0));
359 let captured_waker = Arc::new(Mutex::new(None::<Waker>));
360 let future_polls = Arc::clone(&polls);
361 let future_waker = Arc::clone(&captured_waker);
362 let mut handle = Box::pin(executor.spawn(futures::future::poll_fn(move |context| {
363 future_polls.fetch_add(1, Ordering::SeqCst);
364 *future_waker
365 .lock()
366 .expect("captured-waker mutex must remain available") =
367 Some(context.waker().clone());
368 Poll::Ready(5usize)
369 })));
370
371 executor.process_pending_tasks();
372 assert_eq!(executor.stats().tasks_completed, 1);
373 assert_eq!(polls.load(Ordering::SeqCst), 1);
374 assert!(executor.run_queue.is_empty());
375
376 captured_waker
377 .lock()
378 .expect("captured-waker mutex must remain available")
379 .take()
380 .expect("the completed future must capture its executor waker")
381 .wake();
382 assert!(
383 executor.run_queue.is_empty(),
384 "a stale wake must not requeue a completed task"
385 );
386
387 executor.process_pending_tasks();
388 assert_eq!(polls.load(Ordering::SeqCst), 1);
389 assert_eq!(executor.stats().tasks_completed, 1);
390
391 let handle_waker = futures::task::noop_waker();
392 let mut context = Context::from_waker(&handle_waker);
393 assert!(matches!(handle.as_mut().poll(&mut context), Poll::Ready(5)));
394 }
395
396 #[test]
397 fn completion_under_lock_blocks_a_concurrent_polling_thread() {
398 use std::sync::atomic::AtomicUsize;
399 use std::sync::{Arc, Barrier};
400
401 let executor = Arc::new(AsyncExecutor::new().unwrap());
416 let polls = Arc::new(AtomicUsize::new(0));
417
418 let poll_counter = Arc::clone(&polls);
421 let _handle = executor.spawn(async move {
422 assert_eq!(
423 poll_counter.fetch_add(1, Ordering::SeqCst),
424 0,
425 "future must never be polled after completion"
426 );
427 });
428
429 let task = executor
432 .run_queue
433 .try_dequeue()
434 .expect("spawned task must be queued");
435 let guard = task.future_lock.lock().unwrap();
436 executor.run_queue.enqueue(Arc::clone(&task));
437
438 let barrier = Arc::new(Barrier::new(2));
439 let poller = {
440 let executor = Arc::clone(&executor);
441 let barrier = Arc::clone(&barrier);
442 std::thread::spawn(move || {
443 barrier.wait();
444 executor.process_pending_tasks();
446 })
447 };
448
449 barrier.wait();
450
451 task.completed.store(true, Ordering::Release);
456 drop(guard);
457
458 poller.join().expect("polling thread must not panic");
459
460 assert_eq!(
461 polls.load(Ordering::SeqCst),
462 0,
463 "a task completed while another thread waited on its lock must not be polled"
464 );
465 }
466
467 #[test]
468 fn test_priority_scheduling() {
469 let executor = AsyncExecutor::new().unwrap();
470
471 let _high_priority = executor.spawn_with_priority(async { "high" }, Priority::High);
472 let _normal_priority = executor.spawn_with_priority(async { "normal" }, Priority::Normal);
473 let _low_priority = executor.spawn_with_priority(async { "low" }, Priority::Low);
474
475 let stats = executor.stats();
476 assert_eq!(stats.tasks_spawned, 3);
477 assert_eq!(stats.tasks_pending, 3);
478 }
479}