use crate::journal::frame::SaturatingFrom;
use std::future::Future;
use std::num::{NonZeroU32, NonZeroUsize};
use std::sync::Arc;
use std::time::Duration;
use crate::rt::clock::{Clock, TimeSource};
use crate::rt::time;
use crate::rt::time::Deadline;
use super::{FlowError, MAX_ATTEMPTS, MAX_BACKOFF, MAX_IN_FLIGHT, Scope, StepKey};
#[doc(hidden)]
pub fn at_most(limit: usize) -> Result<NonZeroUsize, FlowError> {
NonZeroUsize::new(limit)
.filter(|bound| bound.get() <= MAX_IN_FLIGHT)
.ok_or(FlowError::InvalidBound {
what: "each: at most N at once",
value: u64::saturating_from(limit),
max: u64::saturating_from(MAX_IN_FLIGHT),
})
}
#[doc(hidden)]
pub fn attempts(count: u32) -> Result<NonZeroU32, FlowError> {
NonZeroU32::new(count)
.filter(|budget| budget.get() <= MAX_ATTEMPTS)
.ok_or(FlowError::InvalidBound {
what: "retry: up to N times",
value: u64::from(count),
max: u64::from(MAX_ATTEMPTS),
})
}
#[inline(never)]
fn note_refusal(site: &str, error: &FlowError) {
lgwks_std::trace::debug!(site, ?error, "a step was refused before it ran");
}
#[inline(always)]
fn refuse_step<T>(site: &str, error: FlowError) -> Result<T, FlowError> {
note_refusal(site, &error);
Err(error)
}
pub async fn within<T, Fut>(
scope: &Scope,
step: &str,
limit: Duration,
body: Fut,
) -> Result<T, FlowError>
where
Fut: Future<Output = Result<T, FlowError>>,
{
within_on(scope.clock(), scope, step, limit, body).await
}
pub async fn within_on<T, Fut>(
clock: &Clock,
scope: &Scope,
step: &str,
limit: Duration,
body: Fut,
) -> Result<T, FlowError>
where
Fut: Future<Output = Result<T, FlowError>>,
{
let deadline = Deadline::after(clock, limit);
let at = scope.join(step);
let timed_out = || FlowError::TimedOut {
at: Arc::clone(&at),
after: limit,
};
if deadline.is_exhausted() {
return refuse_step("within_on", timed_out());
}
let logical = Box::pin(settle_logical_bound(
&deadline,
clock.source() == TimeSource::Wall,
Arc::clone(&at),
limit,
body,
));
let watchdog = Box::pin(time::sleep(limit));
let stopped = Box::pin(scope.token().cancelled());
lgwks_deps::tokio::select! {
biased;
finished = logical => finished,
() = watchdog => Err(timed_out()),
() = stopped => Err(FlowError::Cancelled { at: Arc::clone(&at) }),
}
}
async fn settle_logical_bound<T, Fut>(
deadline: &Deadline<'_>,
wall: bool,
at: Arc<str>,
limit: Duration,
body: Fut,
) -> Result<T, FlowError>
where
Fut: Future<Output = Result<T, FlowError>>,
{
const LOGICAL_POLL: Duration = Duration::from_millis(1);
if wall {
return body.await;
}
let mut body = std::pin::pin!(body);
let refused = || FlowError::TimedOut {
at: Arc::clone(&at),
after: limit,
};
loop {
if deadline.is_exhausted() {
return refuse_step("settle_logical_bound", refused());
}
lgwks_deps::tokio::select! {
biased;
() = time::sleep(LOGICAL_POLL) => {}
finished = &mut body => {
if deadline.is_exhausted() {
return Err(refused());
}
return finished;
}
}
}
}
pub async fn retry<T, F, Fut>(
scope: &Scope,
step: &str,
attempts: NonZeroU32,
backoff: Duration,
mut body: F,
) -> Result<T, FlowError>
where
F: FnMut(Scope, u32) -> Fut,
Fut: Future<Output = Result<T, FlowError>>,
{
let here = scope.enter_sharing(step)?;
here.policy().first_attempt();
let mut attempt: u32 = 1;
loop {
here.checkpoint()?;
let error = match body(here.clone(), attempt).await {
Ok(value) => return Ok(value),
Err(error) => error.located(&here),
};
if !error.is_retryable() {
let refusal = Err(error);
lgwks_std::trace::debug!(error = ?refusal.as_ref().err(), "retry: returning an error to the caller");
return refusal;
}
if attempt >= attempts.get() {
let refusal = Err(FlowError::Exhausted {
at: Arc::clone(here.shared_path()),
attempts: attempt,
last: Box::new(error),
});
lgwks_std::trace::debug!(error = ?refusal.as_ref().err(), "retry: returning an error to the caller");
return refusal;
}
if !here.policy().take_retry() {
let refusal = Err(FlowError::Throttled {
at: Arc::clone(here.shared_path()),
attempts: attempt,
last: Box::new(error),
});
lgwks_std::trace::debug!(error = ?refusal.as_ref().err(), "retry: returning an error to the caller");
return refusal;
}
let delay = backoff_delay(backoff, attempt, &here.key());
if !delay.is_zero()
&& here
.token()
.run_until_cancelled(time::sleep(delay))
.await
.is_none()
{
let refusal = Err(FlowError::Cancelled {
at: Arc::clone(here.shared_path()),
});
lgwks_std::trace::debug!(error = ?refusal.as_ref().err(), "retry: returning an error to the caller");
return refusal;
}
attempt = attempt.saturating_add(1);
}
}
fn backoff_delay(base: Duration, attempt: u32, key: &StepKey) -> Duration {
if base.is_zero() {
return Duration::ZERO;
}
let shift = attempt.saturating_sub(1).min(16);
let mut doubled = base;
for _ in 0..shift {
doubled = doubled.saturating_mul(2);
}
let grown = doubled.min(MAX_BACKOFF);
let jitter_byte = key
.as_bytes()
.get(usize::saturating_from(attempt) & 31)
.copied()
.map_or(0, u32::from);
let scaled = grown.as_nanos().saturating_mul(u128::from(jitter_byte));
let jitter = Duration::from_nanos(u64::saturating_from(scaled.div_euclid(512)));
grown.saturating_add(jitter)
}