use std::{
cell::{Cell, RefCell},
future::Future,
pin::Pin,
rc::Rc,
task::{Context, Poll, Waker},
};
use compio::runtime::{CancelToken, JoinHandle, spawn};
use log::warn;
use whasher::HashMap;
use super::task_type::{TaskPlacementCategory, TaskType};
struct DoneEvent {
done: Cell<bool>,
waiters: RefCell<Vec<Waker>>,
}
impl DoneEvent {
fn new() -> Self {
Self {
done: Cell::new(false),
waiters: RefCell::new(Vec::new()),
}
}
fn notify_done(&self) {
self.done.set(true);
for w in self.waiters.borrow_mut().drain(..) {
w.wake();
}
}
fn completed(&self) -> bool {
self.done.get()
}
fn poll_wait(&self, cx: &mut Context<'_>) -> Poll<()> {
if self.done.get() {
return Poll::Ready(());
}
let mut waiters = self.waiters.borrow_mut();
if self.done.get() {
return Poll::Ready(());
}
if !waiters.iter().any(|w| w.will_wake(cx.waker())) {
waiters.push(cx.waker().clone());
}
Poll::Pending
}
}
struct DoneWait<'a> {
event: &'a DoneEvent,
}
impl Future for DoneWait<'_> {
type Output = ();
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<()> {
self.event.poll_wait(cx)
}
}
struct TaskState {
cts: CancelToken,
task: RefCell<Option<JoinHandle<()>>>,
done: DoneEvent,
}
impl TaskState {
fn wait_done(&self) -> DoneWait<'_> {
DoneWait { event: &self.done }
}
}
type Registry = Rc<RefCell<HashMap<TaskType, Rc<TaskState>>>>;
#[derive(Clone)]
pub struct TaskManager {
cts: CancelToken,
registry: Registry,
}
impl TaskManager {
pub fn new() -> Self {
Self {
cts: CancelToken::new(),
registry: Rc::new(RefCell::new(HashMap::default())),
}
}
#[must_use]
pub fn is_running(&self, task_type: TaskType) -> bool {
self
.registry
.borrow()
.get(&task_type)
.is_some_and(|state| !state.done.completed())
}
#[must_use]
pub fn is_registered(&self, task_type: TaskType) -> bool {
self.registry.borrow().contains_key(&task_type)
}
pub fn register_and_run<F, Fut>(
&self,
task_type: TaskType,
task_factory: F,
cleanup_on_completion: bool,
) -> bool
where
F: FnOnce(CancelToken) -> Fut,
Fut: Future<Output = ()> + 'static,
{
let state = Rc::new(TaskState {
cts: CancelToken::new(),
task: RefCell::new(None),
done: DoneEvent::new(),
});
{
let mut registry = self.registry.borrow_mut();
if registry.contains_key(&task_type) {
warn!("{task_type:?} already registered!");
return false;
}
let fut = task_factory(state.cts.clone());
let task_state = Rc::clone(&state);
let registry_for_cleanup = Rc::clone(&self.registry);
let task = spawn(async move {
fut.await;
task_state.done.notify_done();
if cleanup_on_completion {
registry_for_cleanup.borrow_mut().remove(&task_type);
}
});
state.task.borrow_mut().replace(task);
registry.insert(task_type, state);
}
true
}
pub async fn cancel_async(&self, task_type: TaskType) {
let state = self.registry.borrow_mut().remove(&task_type);
if let Some(state) = state {
state.cts.clone().cancel();
state.wait_done().await;
drop(state.task.borrow_mut().take());
}
}
pub async fn cancel_category_async(&self, task_placement_category: TaskPlacementCategory) {
for task_type in TaskType::get_task_types(task_placement_category) {
self.cancel_async(task_type).await;
}
}
pub async fn wait_async(&self, task_type: TaskType) -> bool {
let Some(state) = self.registry.borrow().get(&task_type).cloned() else {
return false;
};
state.wait_done().await;
true
}
}
impl Default for TaskManager {
fn default() -> Self {
Self::new()
}
}
impl Drop for TaskManager {
fn drop(&mut self) {
self.cts.clone().cancel();
for (_, state) in self.registry.borrow_mut().drain() {
state.cts.clone().cancel();
}
}
}
#[cfg(test)]
mod tests {
use std::{rc::Rc, time::Duration};
use compio::{runtime::Runtime, time::sleep};
use parking_lot::Mutex;
use super::{TaskManager, TaskType};
use crate::taskmanager::task_type::TaskPlacementCategory;
#[test]
fn register_and_run_with_cleanup() {
let rt = Runtime::new().unwrap();
rt.block_on(async {
let ran = Rc::new(Mutex::new(false));
let ran2 = Rc::clone(&ran);
let tm = TaskManager::new();
assert!(tm.register_and_run(
TaskType::CommitTask,
move |_| async move {
*ran2.lock() = true;
},
true
));
assert!(tm.is_registered(TaskType::CommitTask));
assert!(tm.is_running(TaskType::CommitTask));
assert!(!tm.register_and_run(TaskType::CommitTask, |_| async {}, false));
tm.wait_async(TaskType::CommitTask).await;
assert!(*ran.lock());
assert!(!tm.is_registered(TaskType::CommitTask));
assert!(!tm.is_running(TaskType::CommitTask));
});
}
#[test]
fn register_without_cleanup_keeps_entry() {
let rt = Runtime::new().unwrap();
rt.block_on(async {
let tm = TaskManager::new();
tm.register_and_run(TaskType::CompactionTask, |_| async {}, false);
tm.wait_async(TaskType::CompactionTask).await;
assert!(tm.is_registered(TaskType::CompactionTask));
assert!(!tm.is_running(TaskType::CompactionTask));
});
}
#[test]
fn cancel_async_stops_long_running_task() {
let rt = Runtime::new().unwrap();
rt.block_on(async {
let tm = TaskManager::new();
let cancelled = Rc::new(Mutex::new(false));
let cancelled2 = Rc::clone(&cancelled);
tm.register_and_run(
TaskType::ExpiredKeyDeletionTask,
move |token| async move {
token.wait().await;
*cancelled2.lock() = true;
},
false,
);
assert!(tm.is_running(TaskType::ExpiredKeyDeletionTask));
tm.cancel_async(TaskType::ExpiredKeyDeletionTask).await;
assert!(*cancelled.lock());
assert!(!tm.is_registered(TaskType::ExpiredKeyDeletionTask));
});
}
#[test]
fn cancel_category_cancels_matching_placement_only() {
let rt = Runtime::new().unwrap();
rt.block_on(async {
let tm = TaskManager::new();
tm.register_and_run(
TaskType::AofSizeLimitTask,
|t| async move { t.wait().await },
false,
);
tm.register_and_run(
TaskType::VectorReplicationReplayTask,
|t| async move { t.wait().await },
false,
);
tm.cancel_category_async(TaskPlacementCategory::PRIMARY)
.await;
assert!(!tm.is_registered(TaskType::AofSizeLimitTask));
assert!(tm.is_registered(TaskType::VectorReplicationReplayTask));
});
}
#[test]
fn cancel_async_during_sleep_resumes() {
let rt = Runtime::new().unwrap();
rt.block_on(async {
let tm = TaskManager::new();
tm.register_and_run(
TaskType::IndexAutoGrowTask,
|t| async move {
while !t.is_cancelled() {
t.clone().wait().await;
}
},
false,
);
tm.cancel_async(TaskType::IndexAutoGrowTask).await;
assert!(!tm.is_registered(TaskType::IndexAutoGrowTask));
sleep(Duration::from_millis(1)).await;
});
}
#[test]
fn wait_async_unregistered_returns_false() {
let rt = Runtime::new().unwrap();
rt.block_on(async {
let tm = TaskManager::new();
assert!(!tm.wait_async(TaskType::CommitTask).await);
});
}
}