Skip to main content

moirai_executor/hybrid/
manager.rs

1use std::sync::Arc;
2
3use moirai_core::{
4    error::{ExecutorError, ExecutorResult, TaskError},
5    executor::{TaskManager, TaskStats, TaskStatus},
6    task::TaskId,
7};
8
9use super::HybridExecutor;
10use crate::registry::CancelOutcome;
11use crate::schedule::WorkScheduler;
12use crate::task::TaskMetadata;
13
14/// Derive the observable [`TaskStatus`] from registry metadata.
15fn status_of(metadata: &TaskMetadata) -> TaskStatus {
16    if metadata.cancelled {
17        TaskStatus::Cancelled
18    } else if metadata.completed_at.is_some() {
19        TaskStatus::Completed
20    } else if metadata.started_at.is_some() {
21        TaskStatus::Running
22    } else {
23        TaskStatus::Queued
24    }
25}
26
27impl<S: WorkScheduler> TaskManager for HybridExecutor<S> {
28    /// Cooperative cancellation.
29    ///
30    /// Contract: a queued task that has not started skips its body when a
31    /// worker dequeues it — the task completes with `TaskError::Cancelled` and
32    /// status `Cancelled`. A task that already started is **not** preempted; it
33    /// runs to completion and this call still returns `Ok(())` (the request is
34    /// recorded but has no effect). Cancelling an already-completed task is a
35    /// no-op `Ok(())`. An unknown task ID is an error.
36    fn cancel_task(&self, id: TaskId) -> ExecutorResult<()> {
37        match self.task_registry.request_cancel(id.0) {
38            Some(CancelOutcome::Requested | CancelOutcome::AlreadyCompleted) => Ok(()),
39            None => Err(ExecutorError::SpawnFailed(TaskError::InvalidOperation)),
40        }
41    }
42
43    fn task_status(&self, id: TaskId) -> Option<TaskStatus> {
44        self.task_registry
45            .get_metadata(id.0)
46            .map(|metadata| status_of(&metadata))
47    }
48
49    /// Event-driven completion wait.
50    ///
51    /// The future registers a completion waker with the task registry and
52    /// returns `Pending`; `mark_completed`/`mark_cancelled` wake it exactly
53    /// (no polling loop and no thread-blocking sleep). The deadline is checked
54    /// on every poll — at creation, on each completion wake, and on any
55    /// external poll after expiry. No in-scope timer exists (the async timer
56    /// lives in `moirai-async`), so if the task never completes, observing the
57    /// expiry requires the caller to poll after the deadline (e.g. a
58    /// timeout-aware runtime); a completion wake always resolves promptly.
59    fn wait_for_task(
60        &self,
61        id: TaskId,
62        timeout: Option<core::time::Duration>,
63    ) -> impl core::future::Future<Output = ExecutorResult<()>> + Send {
64        let registry = Arc::clone(&self.task_registry);
65        let deadline = timeout.and_then(|timeout| std::time::Instant::now().checked_add(timeout));
66
67        std::future::poll_fn(move |context| {
68            match registry.completion(id.0) {
69                Some(true) => return std::task::Poll::Ready(Ok(())),
70                Some(false) => {}
71                None => {
72                    return std::task::Poll::Ready(Err(ExecutorError::SpawnFailed(
73                        TaskError::InvalidOperation,
74                    )));
75                }
76            }
77            if deadline.is_some_and(|deadline| std::time::Instant::now() >= deadline) {
78                return std::task::Poll::Ready(Err(ExecutorError::SpawnFailed(TaskError::Timeout)));
79            }
80
81            registry.register_waker(id.0, context.waker());
82            // Re-check after registration: completion publishes the timestamp
83            // before taking the waker, so a completion that raced ahead of the
84            // registration is visible here and must not be lost.
85            if registry.is_completed(id.0) {
86                return std::task::Poll::Ready(Ok(()));
87            }
88
89            std::task::Poll::Pending
90        })
91    }
92
93    /// Statistics limited to what the executor actually tracks.
94    ///
95    /// `priority` is the value recorded at spawn; timing fields come from the
96    /// registry's lifecycle timestamps.
97    fn task_stats(&self, id: TaskId) -> Option<TaskStats> {
98        self.task_registry
99            .get_metadata(id.0)
100            .map(|metadata| TaskStats {
101                id,
102                priority: metadata.priority,
103                status: status_of(&metadata),
104                spawn_time: metadata.created_at,
105                start_time: metadata.started_at,
106                completion_time: metadata.completed_at,
107                cpu_time_ns: metadata
108                    .execution_duration()
109                    .map_or(0, |duration| duration.as_nanos() as u64),
110            })
111    }
112}