use std::error::Error;
use std::fmt;
use std::time::{Duration, Instant};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum DeadlineAt {
At(Instant),
Never,
}
#[derive(Debug, Clone, Copy)]
pub struct Deadline {
at: DeadlineAt,
created_at: Instant,
}
impl Deadline {
pub fn after(duration: Duration) -> Self {
let now = Instant::now();
let at = match now.checked_add(duration) {
Some(t) => DeadlineAt::At(t),
None => DeadlineAt::Never,
};
Self { at, created_at: now }
}
pub fn at(instant: Instant) -> Self {
Self {
at: DeadlineAt::At(instant),
created_at: Instant::now(),
}
}
pub fn never() -> Self {
Self {
at: DeadlineAt::Never,
created_at: Instant::now(),
}
}
pub fn is_never(&self) -> bool {
matches!(self.at, DeadlineAt::Never)
}
pub fn instant(&self) -> Option<Instant> {
match self.at {
DeadlineAt::At(i) => Some(i),
DeadlineAt::Never => None,
}
}
pub fn expired(&self) -> bool {
match self.at {
DeadlineAt::Never => false,
DeadlineAt::At(at) => Instant::now() >= at,
}
}
pub fn remaining(&self) -> Duration {
match self.at {
DeadlineAt::Never => Duration::MAX,
DeadlineAt::At(at) => at.saturating_duration_since(Instant::now()),
}
}
pub fn remaining_seconds(&self) -> f64 {
match self.at {
DeadlineAt::Never => f64::INFINITY,
DeadlineAt::At(_) => self.remaining().as_secs_f64(),
}
}
pub fn elapsed(&self) -> Duration {
self.created_at.elapsed()
}
pub fn check_or_err(&self) -> Result<(), DeadlineExceeded> {
match self.at {
DeadlineAt::Never => Ok(()),
DeadlineAt::At(at) => {
let now = Instant::now();
if now >= at {
Err(DeadlineExceeded {
elapsed: now.saturating_duration_since(self.created_at),
})
} else {
Ok(())
}
}
}
}
pub fn intersect(&self, other: &Deadline) -> Deadline {
let tighter = match (self.at, other.at) {
(DeadlineAt::Never, other_at) => other_at,
(self_at, DeadlineAt::Never) => self_at,
(DeadlineAt::At(a), DeadlineAt::At(b)) => DeadlineAt::At(a.min(b)),
};
Deadline {
at: tighter,
created_at: self.created_at,
}
}
pub fn intersect_after(&self, duration: Duration) -> Deadline {
self.intersect(&Deadline::after(duration))
}
}
impl Default for Deadline {
fn default() -> Self {
Self::never()
}
}
impl PartialEq for Deadline {
fn eq(&self, other: &Self) -> bool {
self.at == other.at
}
}
impl Eq for Deadline {}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct DeadlineExceeded {
pub elapsed: Duration,
}
impl DeadlineExceeded {
pub fn elapsed_seconds(&self) -> f64 {
self.elapsed.as_secs_f64()
}
}
impl fmt::Display for DeadlineExceeded {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(
f,
"deadline exceeded: elapsed {:.6}s past the configured cap",
self.elapsed_seconds()
)
}
}
impl Error for DeadlineExceeded {}
#[cfg(feature = "serde")]
mod serde_impl {
use super::{Deadline, DeadlineAt};
use serde::{Deserialize, Deserializer, Serialize, Serializer};
use std::time::Duration;
#[derive(Serialize, Deserialize)]
struct DeadlineSnapshot {
never: bool,
remaining_secs: f64,
}
impl Serialize for Deadline {
fn serialize<S: Serializer>(&self, ser: S) -> Result<S::Ok, S::Error> {
let snap = match self.at {
DeadlineAt::Never => DeadlineSnapshot {
never: true,
remaining_secs: 0.0,
},
DeadlineAt::At(_) => DeadlineSnapshot {
never: false,
remaining_secs: self.remaining().as_secs_f64(),
},
};
snap.serialize(ser)
}
}
impl<'de> Deserialize<'de> for Deadline {
fn deserialize<D: Deserializer<'de>>(de: D) -> Result<Self, D::Error> {
let snap = DeadlineSnapshot::deserialize(de)?;
if snap.never {
return Ok(Deadline::never());
}
let remaining = if snap.remaining_secs.is_finite() && snap.remaining_secs > 0.0 {
Duration::from_secs_f64(snap.remaining_secs)
} else {
Duration::ZERO
};
Ok(Deadline::after(remaining))
}
}
}