use core::fmt;
use crate::rate_limit::RateLimit;
use crate::retry::{MonotonicDuration, MonotonicInstant};
use super::progress::{ProgressError, ProgressTracker};
use super::{
PollBackoff, PollContext, PollControl, PollRequestStep, ProgressObservation, ProgressPolicy,
ProviderTimeObservation,
};
#[derive(Eq, PartialEq)]
pub enum ActionUpdate<E> {
Running,
Success,
Failed(E),
}
impl<E> fmt::Debug for ActionUpdate<E> {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Running => formatter.write_str("Running"),
Self::Success => formatter.write_str("Success"),
Self::Failed(_) => formatter.write_str("Failed([redacted])"),
}
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum ActionPollLimitsError {
Zero,
DelayExceedsCumulative,
DelayExceedsElapsed,
}
impl_static_error!(ActionPollLimitsError,
Self::Zero => "action polling limits must be nonzero",
Self::DelayExceedsCumulative => "maximum poll delay exceeds the cumulative budget",
Self::DelayExceedsElapsed => "maximum poll delay reaches the elapsed budget",
);
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct ActionPollLimits {
max_observations: u32,
max_delay: MonotonicDuration,
max_cumulative_delay: MonotonicDuration,
max_elapsed: MonotonicDuration,
}
impl ActionPollLimits {
pub const fn new(
max_observations: u32,
max_delay: MonotonicDuration,
max_cumulative_delay: MonotonicDuration,
max_elapsed: MonotonicDuration,
) -> Result<Self, ActionPollLimitsError> {
if max_observations == 0
|| max_delay.get() == 0
|| max_cumulative_delay.get() == 0
|| max_elapsed.get() == 0
{
return Err(ActionPollLimitsError::Zero);
}
if max_delay.get() > max_cumulative_delay.get() {
return Err(ActionPollLimitsError::DelayExceedsCumulative);
}
if max_delay.get() >= max_elapsed.get() {
return Err(ActionPollLimitsError::DelayExceedsElapsed);
}
Ok(Self {
max_observations,
max_delay,
max_cumulative_delay,
max_elapsed,
})
}
#[must_use]
pub const fn max_observations(self) -> u32 {
self.max_observations
}
#[must_use]
pub const fn max_delay(self) -> MonotonicDuration {
self.max_delay
}
#[must_use]
pub const fn max_cumulative_delay(self) -> MonotonicDuration {
self.max_cumulative_delay
}
#[must_use]
pub const fn max_elapsed(self) -> MonotonicDuration {
self.max_elapsed
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum ActionPollError {
ResponsePending,
UnexpectedObservation,
Terminal,
MonotonicRollback,
TimeOverflow,
ObservationLimitExceeded,
InvalidProgress,
ProgressRegressed,
ProgressResetForbidden,
ProgressResetLimitExceeded,
ZeroDelay,
DelayLimitExceeded,
CumulativeDelayExceeded,
ElapsedBudgetExceeded,
}
impl_static_error!(ActionPollError,
Self::ResponsePending => "an action response is still pending",
Self::UnexpectedObservation => "action observation has no admitted request",
Self::Terminal => "action polling already reached a terminal state",
Self::MonotonicRollback => "action polling monotonic time moved backwards",
Self::TimeOverflow => "action polling monotonic time overflowed",
Self::ObservationLimitExceeded => "action polling observation limit was exceeded",
Self::InvalidProgress => "action progress exceeds 100",
Self::ProgressRegressed => "action progress moved backwards",
Self::ProgressResetForbidden => "action progress reset is forbidden",
Self::ProgressResetLimitExceeded => "action progress reset limit was exceeded",
Self::ZeroDelay => "action poll backoff requested a zero delay",
Self::DelayLimitExceeded => "action poll backoff exceeded the delay limit",
Self::CumulativeDelayExceeded => "action polling cumulative delay was exceeded",
Self::ElapsedBudgetExceeded => "action polling elapsed budget would be exceeded",
);
#[derive(Eq, PartialEq)]
pub enum ActionObserveError<E> {
Driver(ActionPollError),
Backoff(E),
}
impl<E> fmt::Debug for ActionObserveError<E> {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Driver(error) => formatter.debug_tuple("Driver").field(error).finish(),
Self::Backoff(_) => formatter.write_str("Backoff([redacted])"),
}
}
}
impl<E> fmt::Display for ActionObserveError<E> {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Driver(error) => fmt::Display::fmt(error, formatter),
Self::Backoff(_) => formatter.write_str("action poll backoff policy failed"),
}
}
}
impl<E> core::error::Error for ActionObserveError<E> {}
#[derive(Eq, PartialEq)]
pub enum ActionPollStep<E> {
Delay(MonotonicDuration),
Complete,
Failed(E),
}
impl<E> fmt::Debug for ActionPollStep<E> {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Delay(delay) => formatter.debug_tuple("Delay").field(delay).finish(),
Self::Complete => formatter.write_str("Complete"),
Self::Failed(_) => formatter.write_str("Failed([redacted])"),
}
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
enum PollPhase {
Ready,
AwaitingResponse,
Delayed(MonotonicInstant),
Terminal,
}
pub struct ActionPoller {
limits: ActionPollLimits,
progress: ProgressTracker,
started: MonotonicInstant,
last_observed: MonotonicInstant,
observations: u32,
cumulative_delay: u64,
phase: PollPhase,
}
impl ActionPoller {
#[must_use]
pub const fn new(
limits: ActionPollLimits,
progress_policy: ProgressPolicy,
started: MonotonicInstant,
) -> Self {
Self {
limits,
progress: ProgressTracker::new(progress_policy),
started,
last_observed: started,
observations: 0,
cumulative_delay: 0,
phase: PollPhase::Ready,
}
}
#[must_use]
pub const fn observations(&self) -> u32 {
self.observations
}
#[must_use]
pub const fn cumulative_delay(&self) -> MonotonicDuration {
MonotonicDuration::new(self.cumulative_delay)
}
#[must_use]
pub const fn is_terminal(&self) -> bool {
matches!(self.phase, PollPhase::Terminal)
}
pub fn next_request(
&mut self,
control: PollControl,
now: MonotonicInstant,
) -> Result<PollRequestStep, ActionPollError> {
if self.is_terminal() {
return Err(ActionPollError::Terminal);
}
if now < self.last_observed {
self.phase = PollPhase::Terminal;
return Err(ActionPollError::MonotonicRollback);
}
self.last_observed = now;
if self.elapsed_exhausted(now)? {
self.phase = PollPhase::Terminal;
return Ok(PollRequestStep::TimedOut);
}
if control == PollControl::Cancel {
self.phase = PollPhase::Terminal;
return Ok(PollRequestStep::Cancelled);
}
match self.phase {
PollPhase::AwaitingResponse => Err(ActionPollError::ResponsePending),
PollPhase::Delayed(not_before) if now < not_before => {
let remaining = not_before
.checked_duration_since(now)
.ok_or(ActionPollError::MonotonicRollback)?;
Ok(PollRequestStep::Delay(remaining))
}
PollPhase::Ready | PollPhase::Delayed(_) => {
self.phase = PollPhase::AwaitingResponse;
Ok(PollRequestStep::Request)
}
PollPhase::Terminal => Err(ActionPollError::Terminal),
}
}
pub fn observe<E, B>(
&mut self,
update: ActionUpdate<E>,
progress: ProgressObservation,
rate_limit: Option<RateLimit>,
provider_time: ProviderTimeObservation,
now: MonotonicInstant,
backoff: &mut B,
) -> Result<ActionPollStep<E>, ActionObserveError<B::Error>>
where
B: PollBackoff,
{
if self.phase != PollPhase::AwaitingResponse {
return Err(ActionObserveError::Driver(if self.is_terminal() {
ActionPollError::Terminal
} else {
ActionPollError::UnexpectedObservation
}));
}
self.validate_observation_time(now)
.map_err(ActionObserveError::Driver)?;
let observations = self.observations.checked_add(1).ok_or_else(|| {
self.phase = PollPhase::Terminal;
ActionObserveError::Driver(ActionPollError::ObservationLimitExceeded)
})?;
self.observations = observations;
match update {
ActionUpdate::Success => {
self.phase = PollPhase::Terminal;
Ok(ActionPollStep::Complete)
}
ActionUpdate::Failed(error) => {
self.phase = PollPhase::Terminal;
Ok(ActionPollStep::Failed(error))
}
ActionUpdate::Running => self.observe_running(
observations,
progress,
rate_limit,
provider_time,
now,
backoff,
),
}
}
fn observe_running<E, B>(
&mut self,
observations: u32,
progress: ProgressObservation,
rate_limit: Option<RateLimit>,
provider_time: ProviderTimeObservation,
now: MonotonicInstant,
backoff: &mut B,
) -> Result<ActionPollStep<E>, ActionObserveError<B::Error>>
where
B: PollBackoff,
{
if observations >= self.limits.max_observations {
self.phase = PollPhase::Terminal;
return Err(ActionObserveError::Driver(
ActionPollError::ObservationLimitExceeded,
));
}
let progress_change = self.progress.observe(progress).map_err(|error| {
self.phase = PollPhase::Terminal;
ActionObserveError::Driver(map_progress_error(error))
})?;
let context = PollContext {
observation: observations,
progress,
progress_change,
rate_limit,
provider_time,
};
let delay = backoff.delay(context).map_err(|error| {
self.phase = PollPhase::Terminal;
ActionObserveError::Backoff(error)
})?;
self.schedule(delay, now)
.map_err(ActionObserveError::Driver)?;
Ok(ActionPollStep::Delay(delay))
}
fn schedule(
&mut self,
delay: MonotonicDuration,
now: MonotonicInstant,
) -> Result<(), ActionPollError> {
if delay.get() == 0 {
self.phase = PollPhase::Terminal;
return Err(ActionPollError::ZeroDelay);
}
if delay > self.limits.max_delay {
self.phase = PollPhase::Terminal;
return Err(ActionPollError::DelayLimitExceeded);
}
let Some(cumulative) = self.cumulative_delay.checked_add(delay.get()) else {
self.phase = PollPhase::Terminal;
return Err(ActionPollError::CumulativeDelayExceeded);
};
if cumulative > self.limits.max_cumulative_delay.get() {
self.phase = PollPhase::Terminal;
return Err(ActionPollError::CumulativeDelayExceeded);
}
let Some(not_before) = now.checked_add(delay) else {
self.phase = PollPhase::Terminal;
return Err(ActionPollError::TimeOverflow);
};
let Some(elapsed) = not_before.checked_duration_since(self.started) else {
self.phase = PollPhase::Terminal;
return Err(ActionPollError::MonotonicRollback);
};
if elapsed >= self.limits.max_elapsed {
self.phase = PollPhase::Terminal;
return Err(ActionPollError::ElapsedBudgetExceeded);
}
self.cumulative_delay = cumulative;
self.phase = PollPhase::Delayed(not_before);
Ok(())
}
fn validate_observation_time(&mut self, now: MonotonicInstant) -> Result<(), ActionPollError> {
if now < self.last_observed {
self.phase = PollPhase::Terminal;
return Err(ActionPollError::MonotonicRollback);
}
self.last_observed = now;
if self.elapsed_exhausted(now)? {
self.phase = PollPhase::Terminal;
return Err(ActionPollError::ElapsedBudgetExceeded);
}
Ok(())
}
fn elapsed_exhausted(&self, now: MonotonicInstant) -> Result<bool, ActionPollError> {
let elapsed = now
.checked_duration_since(self.started)
.ok_or(ActionPollError::MonotonicRollback)?;
Ok(elapsed >= self.limits.max_elapsed)
}
}
impl fmt::Debug for ActionPoller {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("ActionPoller")
.field("limits", &self.limits)
.field("observations", &self.observations)
.field("cumulative_delay", &self.cumulative_delay())
.field("phase", &self.phase)
.finish_non_exhaustive()
}
}
fn map_progress_error(error: ProgressError) -> ActionPollError {
match error {
ProgressError::Invalid => ActionPollError::InvalidProgress,
ProgressError::Regressed => ActionPollError::ProgressRegressed,
ProgressError::ResetForbidden => ActionPollError::ProgressResetForbidden,
ProgressError::ResetLimit => ActionPollError::ProgressResetLimitExceeded,
}
}