use std::{
num::NonZeroU32,
sync::Arc,
time::{Duration, Instant},
};
use tokio::sync::Semaphore;
use tokio_util::sync::CancellationToken;
use crate::{
TaskError,
core::runner::run_once,
error::SharedError,
events::{Bus, Event, EventKind},
identity::TaskId,
policies::{BackoffPolicy, RestartPolicy},
reasons,
tasks::Task,
};
const IMMEDIATE_RESTART_FLOOR: Duration = Duration::from_millis(1);
fn floored_interval(interval: Duration, elapsed: Duration) -> Duration {
interval.max(IMMEDIATE_RESTART_FLOOR.saturating_sub(elapsed))
}
#[derive(Debug, Clone)]
pub(crate) enum ActorExitReason {
Completed,
Exhausted {
reason: Arc<str>,
exit_code: Option<i32>,
source: Option<SharedError>,
},
Canceled,
Fatal {
reason: Arc<str>,
exit_code: Option<i32>,
source: Option<SharedError>,
},
}
#[derive(Clone)]
pub(crate) struct TaskActorParams {
pub(crate) restart: RestartPolicy,
pub(crate) backoff: BackoffPolicy,
pub(crate) timeout: Option<Duration>,
pub(crate) max_retries: Option<NonZeroU32>,
}
pub(crate) struct TaskActor {
id: TaskId,
name: Arc<str>,
task: Arc<dyn Task>,
params: TaskActorParams,
bus: Bus,
semaphore: Option<Arc<Semaphore>>,
}
impl TaskActor {
pub(crate) fn new(
bus: Bus,
name: Arc<str>,
task: Arc<dyn Task>,
params: TaskActorParams,
semaphore: Option<Arc<Semaphore>>,
id: TaskId,
) -> Self {
Self {
id,
name,
task,
params,
bus,
semaphore,
}
}
pub(crate) async fn run(self, runtime_token: CancellationToken) -> ActorExitReason {
let task_name: Arc<str> = self.name.clone();
let id = self.id;
let mut attempt: u32 = 0;
let mut backoff_attempt: u32 = 0;
loop {
if runtime_token.is_cancelled() {
return ActorExitReason::Canceled;
}
let permit = match &self.semaphore {
Some(sem) => {
let fut = sem.clone().acquire_owned();
tokio::pin!(fut);
tokio::select! {
res = &mut fut => match res {
Ok(p) => Some(p),
Err(_closed) => {
self.bus.publish(
Event::new(EventKind::ActorExhausted)
.with_task(task_name.clone())
.with_id(id)
.with_attempt(attempt)
.with_reason("semaphore_closed"),
);
return ActorExitReason::Canceled;
}
},
_ = runtime_token.cancelled() => {
return ActorExitReason::Canceled;
}
}
}
None => None,
};
if runtime_token.is_cancelled() {
drop(permit);
return ActorExitReason::Canceled;
}
attempt = attempt.saturating_add(1);
self.bus.publish(
Event::new(EventKind::TaskStarting)
.with_task(task_name.clone())
.with_id(id)
.with_attempt(attempt),
);
let attempt_start = Instant::now();
let res = run_once(
self.task.as_ref(),
&runtime_token,
self.params.timeout,
attempt,
id,
&self.bus,
)
.await;
drop(permit);
match res {
Ok(()) => {
backoff_attempt = 0;
match self.params.restart {
RestartPolicy::Always { interval } => {
if let Some(d) = interval {
let delay = floored_interval(d, attempt_start.elapsed());
self.bus.publish(
Event::new(EventKind::BackoffScheduled)
.with_backoff_success()
.with_task(task_name.clone())
.with_id(id)
.with_attempt(attempt)
.with_delay(delay),
);
if !Self::sleep_cancellable(delay, &runtime_token).await {
return ActorExitReason::Canceled;
}
} else {
let elapsed = attempt_start.elapsed();
if elapsed < IMMEDIATE_RESTART_FLOOR {
if !Self::sleep_cancellable(
IMMEDIATE_RESTART_FLOOR - elapsed,
&runtime_token,
)
.await
{
return ActorExitReason::Canceled;
}
} else {
tokio::task::yield_now().await;
}
}
continue;
}
RestartPolicy::OnFailure | RestartPolicy::Never => {
if runtime_token.is_cancelled() {
return ActorExitReason::Canceled;
}
self.bus.publish(
Event::new(EventKind::ActorExhausted)
.with_task(task_name.clone())
.with_id(id)
.with_attempt(attempt)
.with_reason(reasons::POLICY_EXHAUSTED_SUCCESS),
);
return ActorExitReason::Completed;
}
}
}
Err(e) if e.is_fatal() => {
let reason: Arc<str> = Arc::from(e.to_string());
let exit_code = e.exit_code();
let source: Option<SharedError> = e.into_source().map(Arc::from);
let mut ev = Event::new(EventKind::ActorDead)
.with_task(task_name.clone())
.with_id(id)
.with_attempt(attempt)
.with_reason(Arc::clone(&reason));
if let Some(code) = exit_code {
ev = ev.with_exit_code(code);
}
self.bus.publish(ev);
return ActorExitReason::Fatal {
reason,
exit_code,
source,
};
}
Err(TaskError::Canceled) => {
if runtime_token.is_cancelled() {
return ActorExitReason::Canceled;
}
self.bus.publish(
Event::new(EventKind::ActorExhausted)
.with_task(task_name.clone())
.with_id(id)
.with_attempt(attempt)
.with_reason(reasons::TASK_RETURNED_CANCELED),
);
return ActorExitReason::Canceled;
}
Err(e) => {
let policy_allows_retry = matches!(
self.params.restart,
RestartPolicy::OnFailure | RestartPolicy::Always { .. }
);
let error_is_retryable = e.is_retryable();
let retries_exhausted = self
.params
.max_retries
.is_some_and(|max| backoff_attempt >= max.get());
if !(policy_allows_retry && error_is_retryable) || retries_exhausted {
let reason: Arc<str> = if let Some(limit) =
self.params.max_retries.filter(|_| retries_exhausted)
{
Arc::from(format!(
"{}({}/{}): {}",
reasons::MAX_RETRIES_EXCEEDED,
backoff_attempt,
limit.get(),
e
))
} else {
Arc::from(e.to_string())
};
let exit_code = e.exit_code();
let source: Option<SharedError> = e.into_source().map(Arc::from);
let mut ev = Event::new(EventKind::ActorExhausted)
.with_task(task_name.clone())
.with_id(id)
.with_attempt(attempt)
.with_reason(Arc::clone(&reason));
if let Some(code) = exit_code {
ev = ev.with_exit_code(code);
}
self.bus.publish(ev);
return ActorExitReason::Exhausted {
reason,
exit_code,
source,
};
}
let delay = self.params.backoff.delay_for_retry(backoff_attempt);
backoff_attempt = backoff_attempt.saturating_add(1);
self.bus.publish(
Event::new(EventKind::BackoffScheduled)
.with_backoff_failure()
.with_task(task_name.clone())
.with_id(id)
.with_delay(delay)
.with_attempt(attempt)
.with_reason(e.to_string()),
);
if !Self::sleep_cancellable(delay, &runtime_token).await {
return ActorExitReason::Canceled;
}
}
}
}
}
#[inline]
async fn sleep_cancellable(duration: Duration, token: &CancellationToken) -> bool {
let sleep = tokio::time::sleep(duration);
tokio::pin!(sleep);
tokio::select! {
_ = &mut sleep => true,
_ = token.cancelled() => false,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::TaskContext;
use std::future::Future;
use std::pin::Pin;
use std::sync::atomic::{AtomicU32, Ordering};
type BoxFut = Pin<Box<dyn Future<Output = Result<(), TaskError>> + Send + 'static>>;
fn fast_backoff() -> BackoffPolicy {
BackoffPolicy::new(
Duration::from_millis(1),
Duration::from_millis(1),
1.0,
crate::JitterPolicy::None,
)
.expect("valid backoff")
}
fn params(restart: RestartPolicy, max_retries: u32) -> TaskActorParams {
TaskActorParams {
restart,
backoff: fast_backoff(),
timeout: None,
max_retries: NonZeroU32::new(max_retries),
}
}
fn actor(task: Arc<dyn Task>, restart: RestartPolicy, max_retries: u32) -> TaskActor {
let name: Arc<str> = Arc::from(task.name());
TaskActor::new(
Bus::new(16),
name,
Arc::clone(&task),
params(restart, max_retries),
None,
TaskId::next(),
)
}
struct OkTask;
impl Task for OkTask {
fn name(&self) -> &str {
"ok"
}
fn spawn(&self, _ctx: TaskContext) -> BoxFut {
Box::pin(async { Ok(()) })
}
}
struct FailTask;
impl Task for FailTask {
fn name(&self) -> &str {
"fail"
}
fn spawn(&self, _ctx: TaskContext) -> BoxFut {
Box::pin(async { Err(TaskError::fail("boom")) })
}
}
struct FatalTask;
impl Task for FatalTask {
fn name(&self) -> &str {
"fatal"
}
fn spawn(&self, _ctx: TaskContext) -> BoxFut {
Box::pin(async { Err(TaskError::fatal("fatal")) })
}
}
struct CountedTask {
remaining: AtomicU32,
}
impl CountedTask {
fn new(fail_count: u32) -> Self {
Self {
remaining: AtomicU32::new(fail_count),
}
}
}
impl Task for CountedTask {
fn name(&self) -> &str {
"counted"
}
fn spawn(&self, _ctx: TaskContext) -> BoxFut {
let prev = self.remaining.fetch_sub(1, Ordering::SeqCst);
if prev > 0 {
Box::pin(async { Err(TaskError::fail("transient")) })
} else {
Box::pin(async { Ok(()) })
}
}
}
#[tokio::test]
async fn ok_task_returns_completed_under_non_restarting_policies() {
for restart in [RestartPolicy::Never, RestartPolicy::OnFailure] {
let a = actor(Arc::new(OkTask), restart, 0);
let reason = a.run(CancellationToken::new()).await;
assert!(
matches!(reason, ActorExitReason::Completed),
"{restart:?} + Ok task must exit Completed, got {reason:?}"
);
}
}
#[tokio::test]
async fn fatal_error_returns_fatal_with_reason() {
let a = actor(Arc::new(FatalTask), RestartPolicy::OnFailure, 0);
let reason = a.run(CancellationToken::new()).await;
match reason {
ActorExitReason::Fatal {
reason, exit_code, ..
} => {
assert!(
reason.contains("fatal"),
"reason must carry the error: {reason}"
);
assert_eq!(exit_code, None);
}
other => panic!("expected Fatal, got {other:?}"),
}
}
#[tokio::test]
async fn max_retries_exhausted_returns_exhausted_with_reason() {
let a = actor(Arc::new(FailTask), RestartPolicy::OnFailure, 3);
let reason = a.run(CancellationToken::new()).await;
match reason {
ActorExitReason::Exhausted { reason, .. } => {
assert!(
reason.contains("max_retries_exceeded"),
"reason must mention exhausted budget: {reason}"
);
}
other => panic!("expected Exhausted, got {other:?}"),
}
}
#[tokio::test]
async fn cancellation_returns_cancelled() {
let token = CancellationToken::new();
token.cancel();
let a = actor(
Arc::new(OkTask),
RestartPolicy::Always { interval: None },
0,
);
let reason = a.run(token).await;
assert!(matches!(reason, ActorExitReason::Canceled));
}
#[tokio::test]
async fn on_failure_retries_then_succeeds() {
let task = Arc::new(CountedTask::new(2));
let a = actor(task, RestartPolicy::OnFailure, 0);
let reason = a.run(CancellationToken::new()).await;
assert!(matches!(reason, ActorExitReason::Completed));
}
#[tokio::test(start_paused = true)]
async fn always_none_instant_ok_is_rate_limited() {
use std::sync::atomic::{AtomicU32, Ordering};
struct Counting(Arc<AtomicU32>);
impl Task for Counting {
fn name(&self) -> &str {
"spin"
}
fn spawn(&self, _ctx: TaskContext) -> BoxFut {
self.0.fetch_add(1, Ordering::Relaxed);
Box::pin(async { Ok(()) })
}
}
let counter = Arc::new(AtomicU32::new(0));
let task = Arc::new(Counting(Arc::clone(&counter)));
let a = actor(task, RestartPolicy::Always { interval: None }, 0);
let token = CancellationToken::new();
let child = token.clone();
let handle = tokio::spawn(async move { a.run(child).await });
tokio::time::sleep(Duration::from_millis(25)).await;
token.cancel();
let _ = handle.await;
let n = counter.load(Ordering::Relaxed);
assert!(
(1..=200).contains(&n),
"Always {{ interval: None }} with an instant-Ok task must be floored, got {n} restarts in 25ms"
);
}
#[test]
fn floored_interval_floors_only_the_idle_portion() {
let floor = IMMEDIATE_RESTART_FLOOR;
assert_eq!(floored_interval(Duration::ZERO, Duration::ZERO), floor);
assert_eq!(floored_interval(floor / 2, Duration::ZERO), floor);
assert_eq!(
floored_interval(Duration::ZERO, floor * 2),
Duration::ZERO,
"a slow attempt must not be additionally delayed"
);
assert_eq!(floored_interval(floor * 10, Duration::ZERO), floor * 10);
assert_eq!(floored_interval(floor * 10, floor * 3), floor * 10);
}
}