moirai_executor/registry/
registry.rs1use std::time::Duration;
2
3use super::super::task::TaskMetadata;
4use super::state::{task_location, TaskState, TaskStateBlock};
5use super::token::TaskLifecycleToken;
6
7#[derive(Debug, Clone, Copy, PartialEq, Eq)]
9pub(crate) enum CancelOutcome {
10 Requested,
12 AlreadyCompleted,
14}
15
16#[derive(Debug)]
18pub struct TaskRegistry {
19 pub(super) blocks: Vec<TaskStateBlock>,
20 pub(super) next_id: u64,
21}
22
23impl TaskRegistry {
24 #[must_use]
26 pub fn new() -> Self {
27 Self {
28 blocks: Vec::new(),
29 next_id: 1,
30 }
31 }
32
33 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 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 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 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 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 #[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 #[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 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 #[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 #[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 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 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}