use std::time::Duration;
use super::super::task::TaskMetadata;
use super::state::{task_location, TaskState, TaskStateBlock};
use super::token::TaskLifecycleToken;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum CancelOutcome {
Requested,
AlreadyCompleted,
}
#[derive(Debug)]
pub struct TaskRegistry {
pub(super) blocks: Vec<TaskStateBlock>,
pub(super) next_id: u64,
}
impl TaskRegistry {
#[must_use]
pub fn new() -> Self {
Self {
blocks: Vec::new(),
next_id: 1,
}
}
pub fn register_task(&mut self) -> u64 {
let id = self.next_id;
self.next_id = self.next_id.saturating_add(1);
self.register_task_with_id(id);
id
}
pub(crate) fn register_next_task(&mut self) -> (u64, TaskLifecycleToken) {
let id = self.next_id;
let lifecycle = self.register_task_with_id(id);
(id, lifecycle)
}
pub(crate) fn register_task_with_id(&mut self, id: u64) -> TaskLifecycleToken {
self.next_id = self.next_id.max(id.saturating_add(1));
let (block_index, slot_index) = task_location(id);
self.ensure_block(block_index);
let block = &self.blocks[block_index];
assert!(
block
.get(slot_index)
.is_none_or(|state| state.is_completed()),
"task ID must not be re-registered while active"
);
TaskLifecycleToken {
state: block.insert(slot_index),
}
}
pub fn mark_started(&self, task_id: u64, worker_id: usize) {
if let Some(state) = self.state(task_id) {
state.mark_started(worker_id);
}
}
pub fn mark_completed(&self, task_id: u64) {
if let Some(state) = self.state(task_id) {
state.mark_completed();
}
}
#[must_use]
pub fn is_completed(&self, task_id: u64) -> bool {
self.state(task_id).is_some_and(TaskState::is_completed)
}
#[must_use]
pub fn get_metadata(&self, task_id: u64) -> Option<TaskMetadata> {
self.state(task_id).map(|state| state.snapshot(task_id))
}
pub fn cleanup_completed(&mut self, older_than: Duration) {
let cutoff = std::time::Instant::now() - older_than;
for block in &self.blocks {
for slot_index in 0..block.len() {
if block
.get(slot_index)
.and_then(TaskState::completed_at)
.is_some_and(|completed| completed <= cutoff)
{
block.clear(slot_index);
}
}
}
while self.blocks.last().is_some_and(TaskStateBlock::is_empty) {
self.blocks.pop();
}
}
#[must_use]
pub fn active_count(&self) -> usize {
self.blocks
.iter()
.flat_map(|block| block.states())
.filter(|state| !state.is_completed())
.count()
}
#[must_use]
pub fn completed_count(&self) -> usize {
self.blocks
.iter()
.flat_map(|block| block.states())
.filter(|state| state.is_completed())
.count()
}
pub(super) fn ensure_block(&mut self, block_index: usize) {
while self.blocks.len() <= block_index {
self.blocks.push(TaskStateBlock::new());
}
}
pub(super) fn state(&self, task_id: u64) -> Option<&TaskState> {
let (block_index, slot_index) = task_location(task_id);
self.blocks.get(block_index)?.get(slot_index)
}
pub(crate) fn request_cancel(&self, task_id: u64) -> Option<CancelOutcome> {
let state = self.state(task_id)?;
if state.is_completed() {
Some(CancelOutcome::AlreadyCompleted)
} else {
state.request_cancel();
Some(CancelOutcome::Requested)
}
}
pub fn register_waker(&self, task_id: u64, waker: &std::task::Waker) -> bool {
if let Some(state) = self.state(task_id) {
let mut guard = state.waker.lock().unwrap();
*guard = Some(waker.clone());
true
} else {
false
}
}
}
impl Default for TaskRegistry {
fn default() -> Self {
Self::new()
}
}