Skip to main content

moirai_executor/registry/
registry.rs

1use std::time::Duration;
2
3use super::super::task::TaskMetadata;
4use super::state::{task_location, TaskState, TaskStateBlock};
5use super::token::TaskLifecycleToken;
6
7/// Outcome of a cooperative cancel request against a registered task.
8#[derive(Debug, Clone, Copy, PartialEq, Eq)]
9pub(crate) enum CancelOutcome {
10    /// The cancel flag was set; the task body is skipped if it has not started.
11    Requested,
12    /// The task already completed; cancelling is a no-op.
13    AlreadyCompleted,
14}
15
16/// Public task registry facade used by executor lifecycle tracking and tests.
17#[derive(Debug)]
18pub struct TaskRegistry {
19    pub(super) blocks: Vec<TaskStateBlock>,
20    pub(super) next_id: u64,
21}
22
23impl TaskRegistry {
24    /// Create a new task registry.
25    #[must_use]
26    pub fn new() -> Self {
27        Self {
28            blocks: Vec::new(),
29            next_id: 1,
30        }
31    }
32
33    /// Register a new task and return its ID.
34    pub fn register_task(&mut self) -> u64 {
35        let id = self.next_id;
36        self.next_id = self.next_id.saturating_add(1);
37        self.register_task_with_id(id);
38        id
39    }
40
41    /// Register a new task and return its ID plus lifecycle mutation token.
42    pub(crate) fn register_next_task(&mut self) -> (u64, TaskLifecycleToken) {
43        let id = self.next_id;
44        let lifecycle = self.register_task_with_id(id);
45        (id, lifecycle)
46    }
47
48    /// Register a task with an externally allocated ID.
49    pub(crate) fn register_task_with_id(&mut self, id: u64) -> TaskLifecycleToken {
50        self.next_id = self.next_id.max(id.saturating_add(1));
51        let (block_index, slot_index) = task_location(id);
52        self.ensure_block(block_index);
53
54        let block = &self.blocks[block_index];
55        assert!(
56            block
57                .get(slot_index)
58                .is_none_or(|state| state.is_completed()),
59            "task ID must not be re-registered while active"
60        );
61
62        TaskLifecycleToken {
63            state: block.insert(slot_index),
64        }
65    }
66
67    /// Mark a task as started.
68    pub fn mark_started(&self, task_id: u64, worker_id: usize) {
69        if let Some(state) = self.state(task_id) {
70            state.mark_started(worker_id);
71        }
72    }
73
74    /// Mark a task as completed.
75    pub fn mark_completed(&self, task_id: u64) {
76        if let Some(state) = self.state(task_id) {
77            state.mark_completed();
78        }
79    }
80
81    /// Check if a task is completed.
82    #[must_use]
83    pub fn is_completed(&self, task_id: u64) -> bool {
84        self.state(task_id).is_some_and(TaskState::is_completed)
85    }
86
87    /// Get task metadata.
88    #[must_use]
89    pub fn get_metadata(&self, task_id: u64) -> Option<TaskMetadata> {
90        self.state(task_id).map(|state| state.snapshot(task_id))
91    }
92
93    /// Remove old completed tasks to prevent retained task metadata growth.
94    pub fn cleanup_completed(&mut self, older_than: Duration) {
95        let cutoff = std::time::Instant::now() - older_than;
96        for block in &self.blocks {
97            for slot_index in 0..block.len() {
98                if block
99                    .get(slot_index)
100                    .and_then(TaskState::completed_at)
101                    .is_some_and(|completed| completed <= cutoff)
102                {
103                    block.clear(slot_index);
104                }
105            }
106        }
107
108        while self.blocks.last().is_some_and(TaskStateBlock::is_empty) {
109            self.blocks.pop();
110        }
111    }
112
113    /// Get count of active tasks.
114    #[must_use]
115    pub fn active_count(&self) -> usize {
116        self.blocks
117            .iter()
118            .flat_map(|block| block.states())
119            .filter(|state| !state.is_completed())
120            .count()
121    }
122
123    /// Get count of completed tasks.
124    #[must_use]
125    pub fn completed_count(&self) -> usize {
126        self.blocks
127            .iter()
128            .flat_map(|block| block.states())
129            .filter(|state| state.is_completed())
130            .count()
131    }
132
133    pub(super) fn ensure_block(&mut self, block_index: usize) {
134        while self.blocks.len() <= block_index {
135            self.blocks.push(TaskStateBlock::new());
136        }
137    }
138
139    pub(super) fn state(&self, task_id: u64) -> Option<&TaskState> {
140        let (block_index, slot_index) = task_location(task_id);
141        self.blocks.get(block_index)?.get(slot_index)
142    }
143
144    /// Request cooperative cancellation of a task.
145    ///
146    /// Returns `None` when the task is unknown. Running tasks are not
147    /// preempted: a task that already started keeps running to completion and
148    /// reports `Requested` here without effect.
149    pub(crate) fn request_cancel(&self, task_id: u64) -> Option<CancelOutcome> {
150        let state = self.state(task_id)?;
151        if state.is_completed() {
152            Some(CancelOutcome::AlreadyCompleted)
153        } else {
154            state.request_cancel();
155            Some(CancelOutcome::Requested)
156        }
157    }
158
159    /// Register a waker to be notified when the task completes.
160    pub fn register_waker(&self, task_id: u64, waker: &std::task::Waker) -> bool {
161        if let Some(state) = self.state(task_id) {
162            let mut guard = state.waker.lock().unwrap();
163            *guard = Some(waker.clone());
164            true
165        } else {
166            false
167        }
168    }
169}
170
171impl Default for TaskRegistry {
172    fn default() -> Self {
173        Self::new()
174    }
175}