1use moirai_core::{Priority, TaskId};
2use moirai_pal::reactor::IoReactor;
3use std::future::Future;
4use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
5use std::sync::Arc;
6use std::task::{Context, Poll, Waker};
7use std::time::Instant;
8
9use crate::executor::handle::AsyncHandle;
10use crate::executor::result_slot::AsyncResultSlot;
11use crate::executor::stats::{AsyncExecutorStats, ExecutorStats};
12use crate::executor::task::{AsyncTask, ErasedTaskFuture};
13use crate::executor::waker::ExecutorWaker;
14
15pub struct AsyncExecutor {
17 reactor: Arc<IoReactor>,
19 run_queue: Arc<moirai_utils::queue::LockFreeQueue<Arc<AsyncTask>>>,
21 stats: AsyncExecutorStats,
23 running: Arc<AtomicBool>,
25 next_task_id: AtomicU64,
27}
28
29impl AsyncExecutor {
30 pub fn new() -> std::io::Result<Self> {
32 let reactor = Arc::new(IoReactor::new()?);
33
34 Ok(Self {
35 reactor,
36 run_queue: Arc::new(moirai_utils::queue::LockFreeQueue::new()),
37 stats: AsyncExecutorStats::default(),
38 running: Arc::new(AtomicBool::new(false)),
39 next_task_id: AtomicU64::new(0),
40 })
41 }
42
43 pub fn spawn<F, T>(&self, future: F) -> AsyncHandle<T>
45 where
46 F: Future<Output = T> + Send + 'static,
47 T: Send + 'static,
48 {
49 self.spawn_with_priority(future, Priority::Normal)
50 }
51
52 pub fn spawn_with_priority<F, T>(&self, future: F, priority: Priority) -> AsyncHandle<T>
54 where
55 F: Future<Output = T> + Send + 'static,
56 T: Send + 'static,
57 {
58 let task_id = TaskId::new(self.next_task_id.fetch_add(1, Ordering::Relaxed));
59 let result_slot = Arc::new(AsyncResultSlot::new());
60 let completion_slot = Arc::clone(&result_slot);
61
62 let wrapped_future = async move {
63 let result = future.await;
64 completion_slot.complete(result);
65 };
66
67 let task = Arc::new(AsyncTask {
68 task_id,
69 future: std::cell::UnsafeCell::new(ErasedTaskFuture::new(wrapped_future)),
70 future_lock: std::sync::Mutex::new(()),
71 is_queued: AtomicBool::new(true),
72 completed: AtomicBool::new(false),
73 priority,
74 created_at: Instant::now(),
75 });
76
77 self.run_queue.enqueue(Arc::clone(&task));
78
79 self.stats.tasks_spawned.fetch_add(1, Ordering::Relaxed);
80 self.stats.tasks_pending.fetch_add(1, Ordering::Relaxed);
81
82 let _ = self.reactor.wake();
83
84 AsyncHandle {
85 task_id,
86 result_slot,
87 }
88 }
89
90 pub fn run(&self) -> std::io::Result<()> {
92 self.running.store(true, Ordering::SeqCst);
93
94 self.reactor.with_active(|| {
95 while self.running.load(Ordering::SeqCst) {
96 self.process_pending_tasks();
97
98 let has_tasks = self.stats.tasks_pending.load(Ordering::Acquire) > 0;
99
100 if !has_tasks {
101 if !self.running.load(Ordering::SeqCst) {
102 break;
103 }
104 self.reactor.run_iteration(None)?;
105 } else {
106 let run_queue_empty = self.run_queue.is_empty();
107 if run_queue_empty {
108 self.reactor.run_iteration(None)?;
109 } else {
110 self.reactor
111 .run_iteration(Some(std::time::Duration::from_millis(0)))?;
112 }
113 }
114 }
115 Ok(())
116 })
117 }
118
119 pub fn stop(&self) -> std::io::Result<()> {
121 self.running.store(false, Ordering::SeqCst);
122 self.reactor.stop()
123 }
124
125 pub(crate) fn process_pending_tasks(&self) {
127 while let Some(task) = self.run_queue.try_dequeue() {
128 task.is_queued.store(false, Ordering::SeqCst);
129
130 let waker = self.create_executor_waker(Arc::clone(&task));
131 let mut context = Context::from_waker(&waker);
132 let task_start = Instant::now();
133
134 let _lock = task.future_lock.lock().unwrap();
150
151 if task.completed.load(Ordering::Acquire) {
152 continue;
153 }
154
155 let future_mut = unsafe { &mut *task.future.get() };
156 match future_mut.poll(&mut context) {
157 std::task::Poll::Ready(()) => {
158 task.completed.store(true, Ordering::Release);
159 self.stats.tasks_completed.fetch_add(1, Ordering::Relaxed);
160 self.stats.tasks_pending.fetch_sub(1, Ordering::Relaxed);
161
162 let execution_time = task_start.elapsed().as_nanos() as u64;
163 self.stats
164 .total_execution_time_ns
165 .fetch_add(execution_time, Ordering::Relaxed);
166 }
167 std::task::Poll::Pending => {}
168 }
169 }
170 }
171
172 fn create_executor_waker(&self, task: Arc<AsyncTask>) -> Waker {
174 let waker = Arc::new(ExecutorWaker {
175 task,
176 run_queue: Arc::clone(&self.run_queue),
177 reactor: Arc::clone(&self.reactor),
178 });
179 Waker::from(waker)
180 }
181
182 pub fn stats(&self) -> ExecutorStats {
184 ExecutorStats {
185 tasks_spawned: self.stats.tasks_spawned.load(Ordering::Relaxed),
186 tasks_completed: self.stats.tasks_completed.load(Ordering::Relaxed),
187 tasks_pending: self.stats.tasks_pending.load(Ordering::Relaxed),
188 total_execution_time_ns: self.stats.total_execution_time_ns.load(Ordering::Relaxed),
189 waker_notifications: self.stats.waker_notifications.load(Ordering::Relaxed),
190 io_operations: self.stats.io_operations.load(Ordering::Relaxed),
191 }
192 }
193
194 pub fn block_on<F, T>(&self, future: F) -> T
196 where
197 F: Future<Output = T> + Send + 'static,
198 T: Send + 'static,
199 {
200 let handle = self.spawn(future);
201 let waker = futures::task::noop_waker();
202 let mut cx = Context::from_waker(&waker);
203 let mut pin_handle = Box::pin(handle);
204
205 self.running.store(true, Ordering::SeqCst);
206
207 loop {
208 self.process_pending_tasks();
209
210 match pin_handle.as_mut().poll(&mut cx) {
211 Poll::Ready(result) => {
212 self.running.store(false, Ordering::SeqCst);
213 return result;
214 }
215 Poll::Pending => {}
216 }
217
218 if self.run_queue.is_empty() {
219 self.reactor.run_iteration(None).ok();
220 } else {
221 self.reactor
222 .run_iteration(Some(std::time::Duration::from_millis(0)))
223 .ok();
224 }
225 }
226 }
227
228 pub fn reactor(&self) -> &IoReactor {
230 &self.reactor
231 }
232}
233
234impl Default for AsyncExecutor {
235 fn default() -> Self {
236 Self::new().expect("Failed to create default AsyncExecutor")
237 }
238}
239
240#[cfg(test)]
241mod tests {
242 use super::*;
243
244 #[test]
245 fn test_native_async_executor_creation() {
246 let executor = AsyncExecutor::new();
247 assert!(executor.is_ok());
248
249 let executor = executor.unwrap();
250 let stats = executor.stats();
251 assert_eq!(stats.tasks_spawned, 0);
252 assert_eq!(stats.tasks_completed, 0);
253 assert_eq!(stats.tasks_pending, 0);
254 }
255
256 #[test]
257 fn test_task_spawning() {
258 let executor = AsyncExecutor::new().unwrap();
259
260 let _handle = executor.spawn(async { 42 });
261
262 let stats = executor.stats();
263 assert_eq!(stats.tasks_spawned, 1);
264 assert_eq!(stats.tasks_pending, 1);
265 }
266
267 #[test]
268 fn test_task_ids_are_unique() {
269 let executor = AsyncExecutor::new().unwrap();
270
271 let first = executor.spawn(async { 1usize });
272 let second = executor.spawn(async { 2usize });
273
274 assert_ne!(first.id(), second.id());
275 }
276
277 #[test]
278 fn test_ready_task_completion_publishes_result() {
279 use std::task::Poll;
280
281 let executor = AsyncExecutor::new().unwrap();
282 let mut handle = Box::pin(executor.spawn(async { 7usize }));
283 let waker = futures::task::noop_waker();
284 let mut context = Context::from_waker(&waker);
285
286 assert!(matches!(handle.as_mut().poll(&mut context), Poll::Pending));
287
288 executor.process_pending_tasks();
289
290 assert_eq!(executor.stats().tasks_completed, 1);
291 assert!(matches!(handle.as_mut().poll(&mut context), Poll::Ready(7)));
292 }
293
294 #[test]
295 fn test_ready_task_completion_wakes_registered_handle() {
296 use futures::task::{waker_ref, ArcWake};
297 use std::sync::atomic::AtomicUsize;
298 use std::task::Poll;
299
300 struct WakeCounter(AtomicUsize);
301
302 impl ArcWake for WakeCounter {
303 fn wake_by_ref(arc_self: &Arc<Self>) {
304 arc_self.0.fetch_add(1, Ordering::SeqCst);
305 }
306 }
307
308 let executor = AsyncExecutor::new().unwrap();
309 let mut handle = Box::pin(executor.spawn(async { 11usize }));
310 let wake_counter = Arc::new(WakeCounter(AtomicUsize::new(0)));
311 let waker = waker_ref(&wake_counter);
312 let mut context = Context::from_waker(&waker);
313
314 assert!(matches!(handle.as_mut().poll(&mut context), Poll::Pending));
315
316 executor.process_pending_tasks();
317
318 assert_eq!(wake_counter.0.load(Ordering::SeqCst), 1);
319 assert!(matches!(
320 handle.as_mut().poll(&mut context),
321 Poll::Ready(11)
322 ));
323 }
324
325 #[test]
326 fn stale_waker_after_completion_does_not_repoll() {
327 let executor = AsyncExecutor::new().unwrap();
332 let _handle = executor.spawn(async { 5usize });
334
335 executor.process_pending_tasks();
337 assert_eq!(executor.stats().tasks_completed, 1);
338
339 let task = executor.run_queue.try_dequeue();
342 assert!(
343 task.is_none(),
344 "completed task must not be on the run queue"
345 );
346
347 executor.process_pending_tasks();
350 assert_eq!(
351 executor.stats().tasks_completed,
352 1,
353 "no re-poll of the completed task"
354 );
355 }
356
357 #[test]
358 fn completion_under_lock_blocks_a_concurrent_polling_thread() {
359 use std::sync::atomic::AtomicUsize;
360 use std::sync::{Arc, Barrier};
361
362 let executor = Arc::new(AsyncExecutor::new().unwrap());
377 let polls = Arc::new(AtomicUsize::new(0));
378
379 let poll_counter = Arc::clone(&polls);
382 let _handle = executor.spawn(async move {
383 assert_eq!(
384 poll_counter.fetch_add(1, Ordering::SeqCst),
385 0,
386 "future must never be polled after completion"
387 );
388 });
389
390 let task = executor
393 .run_queue
394 .try_dequeue()
395 .expect("spawned task must be queued");
396 let guard = task.future_lock.lock().unwrap();
397 executor.run_queue.enqueue(Arc::clone(&task));
398
399 let barrier = Arc::new(Barrier::new(2));
400 let poller = {
401 let executor = Arc::clone(&executor);
402 let barrier = Arc::clone(&barrier);
403 std::thread::spawn(move || {
404 barrier.wait();
405 executor.process_pending_tasks();
407 })
408 };
409
410 barrier.wait();
411
412 task.completed.store(true, Ordering::Release);
417 drop(guard);
418
419 poller.join().expect("polling thread must not panic");
420
421 assert_eq!(
422 polls.load(Ordering::SeqCst),
423 0,
424 "a task completed while another thread waited on its lock must not be polled"
425 );
426 }
427
428 #[test]
429 fn test_priority_scheduling() {
430 let executor = AsyncExecutor::new().unwrap();
431
432 let _high_priority = executor.spawn_with_priority(async { "high" }, Priority::High);
433 let _normal_priority = executor.spawn_with_priority(async { "normal" }, Priority::Normal);
434 let _low_priority = executor.spawn_with_priority(async { "low" }, Priority::Low);
435
436 let stats = executor.stats();
437 assert_eq!(stats.tasks_spawned, 3);
438 assert_eq!(stats.tasks_pending, 3);
439 }
440}