replace AsyncRuntime with simpler CancellableTask

add docstring for AsyncRuntime

task

task
This commit is contained in:
Luna Yao
2026-04-18 17:18:36 +02:00
parent 4512e03d5f
commit b0aae4f1fa
2 changed files with 55 additions and 126 deletions
+1 -1
View File
@@ -62,7 +62,7 @@ futures = { version = "0.3", features = ["bilock", "unstable"] }
tokio = { version = "1", features = ["full"] } tokio = { version = "1", features = ["full"] }
tokio-stream = "0.1" tokio-stream = "0.1"
tokio-util = { version = "0.7.9", features = ["codec", "net", "io"] } tokio-util = { version = "0.7.9", features = ["codec", "net", "io", "rt"] }
async-stream = "0.3.5" async-stream = "0.3.5"
async-trait = "0.1.74" async-trait = "0.1.74"
+54 -125
View File
@@ -1,145 +1,74 @@
use crate::common::scoped_task::ScopedTask;
use derivative::Derivative;
use derive_more::{Deref, DerefMut};
use parking_lot::Mutex;
use std::future::Future; use std::future::Future;
use std::mem::take; use std::pin::Pin;
use std::sync::Arc; use std::task::{Context, Poll};
use std::time::Duration; use std::time::Duration;
use tokio::sync::Notify; use tokio::task::JoinError;
use tokio::task::{AbortHandle, JoinError};
use tokio_util::sync::CancellationToken; use tokio_util::sync::CancellationToken;
use tokio_util::task::AbortOnDropHandle;
#[derive(Derivative, Debug)] #[derive(Debug)]
#[derivative(Default(bound = ""))] pub struct CancellableTask<Output: Send + 'static = ()> {
enum AsyncRuntimeState<R: Send + 'static> { handle: AbortOnDropHandle<Output>,
#[derivative(Default)] token: CancellationToken,
Idle,
Running {
id: tokio::task::Id,
task: ScopedTask<R>,
token: CancellationToken,
},
Stopping(AbortHandle),
} }
#[derive(Derivative, Debug)] impl<Output: Send + 'static> CancellableTask<Output> {
#[derivative(Default(bound = ""))] pub fn token(&self) -> &CancellationToken {
pub struct AsyncRuntimeInner<R: Send + 'static = ()> { &self.token
state: Mutex<AsyncRuntimeState<R>>,
idle: Notify,
}
#[derive(Derivative, Deref, DerefMut)]
#[derivative(Debug = "transparent", Default(bound = ""), Clone(bound = ""))]
pub struct AsyncRuntime<R: Send + 'static = ()>(Arc<AsyncRuntimeInner<R>>);
impl<R: Send + 'static> AsyncRuntime<R> {
pub fn token(&self) -> Option<CancellationToken> {
if let AsyncRuntimeState::Running { token, .. } = &*self.state.lock() {
Some(token.clone())
} else {
None
}
} }
pub fn start<F, Fut>(&self, token: Option<CancellationToken>, factory: F) -> anyhow::Result<()> pub fn new<F>(
token: Option<CancellationToken>,
factory: impl FnOnce(CancellationToken) -> F,
) -> Self
where where
F: FnOnce(CancellationToken) -> Fut, F: Future<Output = Output> + Send + 'static,
Fut: Future<Output = R> + Send + 'static,
{ {
let mut state = self.state.lock();
if !matches!(*state, AsyncRuntimeState::Idle) {
return Err(anyhow::anyhow!("task is already running/stopping"));
}
let token = token.unwrap_or_default(); let token = token.unwrap_or_default();
Self {
let task = { handle: AbortOnDropHandle::new(tokio::spawn(factory(token.clone()))),
let f = factory(token.clone());
let this = (*self).clone();
tokio::spawn(async move {
let result = f.await;
let mut state = this.state.lock();
if let AsyncRuntimeState::Running { id, .. } = &*state
&& *id == tokio::task::id()
{
take(&mut *state);
}
result
})
};
*state = AsyncRuntimeState::Running {
id: task.id(),
task: task.into(),
token, token,
}; }
Ok(())
} }
pub async fn stop(&self, timeout: Option<Duration>) -> Option<Result<R, JoinError>> { pub fn spawn<F>(f: impl FnOnce(CancellationToken) -> F) -> Self
let state = { where
let mut state = self.state.lock(); F: Future<Output = Output> + Send + 'static,
match &*state { {
AsyncRuntimeState::Running { .. } => { Self::new(None, f)
let AsyncRuntimeState::Running { task, token, .. } = take(&mut *state) else {
unreachable!()
};
*state = AsyncRuntimeState::Stopping(task.abort_handle());
Ok((task, token))
}
AsyncRuntimeState::Stopping(_) => Err(self.idle.notified()),
AsyncRuntimeState::Idle => return None,
}
};
let (mut task, token) = match state {
Ok(running) => running,
Err(stopping) => {
stopping.await;
return None;
}
};
token.cancel();
let result = match timeout {
Some(duration) => {
if let Ok(result) = tokio::time::timeout(duration, &mut task).await {
result
} else {
task.abort();
tracing::warn!("task stop timeout after {:?}, aborted", duration);
task.await
}
}
None => task.await,
};
{
let mut state = self.state.lock();
if matches!(*state, AsyncRuntimeState::Stopping(_)) {
*state = AsyncRuntimeState::Idle;
drop(state);
self.idle.notify_waiters();
}
}
Some(result)
} }
pub fn abort(&self) { pub fn child<F>(&self, f: impl FnOnce(CancellationToken) -> F) -> Self
let mut state = self.state.lock(); where
match &*state { F: Future<Output = Output> + Send + 'static,
AsyncRuntimeState::Running { task, .. } => { {
task.abort(); Self::new(Some(self.token.clone()), f)
*state = AsyncRuntimeState::Idle; }
drop(state);
self.idle.notify_waiters(); pub fn with_token<F>(token: CancellationToken, future: F) -> Self
where
F: Future<Output = Output> + Send + 'static,
{
Self::new(Some(token), |_| future)
}
pub async fn stop(mut self, timeout: Option<Duration>) -> Result<Output, JoinError> {
self.token.cancel();
if let Some(timeout) = timeout {
if let Ok(result) = tokio::time::timeout(timeout, &mut self.handle).await {
return result;
} else {
self.handle.abort();
tracing::warn!("task stop timeout after {:?}, aborted", timeout);
} }
AsyncRuntimeState::Stopping(handle) => handle.abort(),
_ => {}
} }
self.handle.await
}
}
impl<Output: Send + 'static> Future for CancellableTask<Output> {
type Output = <AbortOnDropHandle<Output> as Future>::Output;
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
Pin::new(&mut self.handle).poll(cx)
} }
} }