use core::fmt;
#[derive(Debug, Clone, Copy, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[cfg_attr(feature = "serde", serde(try_from = "IntervalRepr"))]
pub struct Interval {
lo: f64,
hi: f64,
}
#[cfg(feature = "serde")]
#[derive(serde::Deserialize)]
struct IntervalRepr {
lo: f64,
hi: f64,
}
#[cfg(feature = "serde")]
impl TryFrom<IntervalRepr> for Interval {
type Error = IntervalError;
fn try_from(r: IntervalRepr) -> Result<Self, IntervalError> {
Interval::new(r.lo, r.hi)
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub enum IntervalError {
Reversed {
lo: f64,
hi: f64,
},
NotANumber,
}
impl fmt::Display for IntervalError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
IntervalError::Reversed { lo, hi } => write!(f, "interval [{lo}, {hi}] is reversed"),
IntervalError::NotANumber => write!(f, "interval end is NaN"),
}
}
}
impl std::error::Error for IntervalError {}
impl Interval {
pub const REAL: Interval = Interval {
lo: f64::NEG_INFINITY,
hi: f64::INFINITY,
};
pub const UNIT: Interval = Interval { lo: 0.0, hi: 1.0 };
pub const TURN: Interval = Interval {
lo: 0.0,
hi: core::f64::consts::TAU,
};
pub fn new(lo: f64, hi: f64) -> Result<Self, IntervalError> {
if lo.is_nan() || hi.is_nan() {
return Err(IntervalError::NotANumber);
}
if lo > hi {
return Err(IntervalError::Reversed { lo, hi });
}
Ok(Interval { lo, hi })
}
pub const fn lo(&self) -> f64 {
self.lo
}
pub const fn hi(&self) -> f64 {
self.hi
}
pub fn length(&self) -> f64 {
self.hi - self.lo
}
pub fn is_bounded(&self) -> bool {
self.lo.is_finite() && self.hi.is_finite()
}
pub fn midpoint(&self) -> f64 {
self.lerp(0.5)
}
pub fn contains(&self, t: f64) -> bool {
self.lo <= t && t <= self.hi
}
pub fn clamp(&self, t: f64) -> f64 {
if t < self.lo {
self.lo
} else if t > self.hi {
self.hi
} else {
t
}
}
pub fn lerp(&self, s: f64) -> f64 {
(1.0 - s) * self.lo + s * self.hi
}
pub fn overlaps(&self, other: &Interval) -> bool {
self.lo <= other.hi && other.lo <= self.hi
}
pub fn intersection(&self, other: &Interval) -> Option<Interval> {
self.overlaps(other).then(|| Interval {
lo: self.lo.max(other.lo),
hi: self.hi.min(other.hi),
})
}
pub fn hull(&self, other: &Interval) -> Interval {
Interval {
lo: self.lo.min(other.lo),
hi: self.hi.max(other.hi),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn iv(lo: f64, hi: f64) -> Interval {
Interval::new(lo, hi).unwrap()
}
#[test]
fn construction_rejects_reversed_and_nan() {
assert_eq!(
Interval::new(2.0, 1.0),
Err(IntervalError::Reversed { lo: 2.0, hi: 1.0 })
);
assert_eq!(Interval::new(f64::NAN, 1.0), Err(IntervalError::NotANumber));
assert!(Interval::new(1.0, 1.0).is_ok());
assert!(Interval::new(f64::NEG_INFINITY, 0.0).is_ok());
}
#[test]
fn clamp_and_contains() {
let i = iv(-1.0, 2.0);
assert!(i.contains(-1.0) && i.contains(2.0) && i.contains(0.5));
assert!(!i.contains(-1.0000001) && !i.contains(f64::NAN));
assert_eq!(i.clamp(-5.0), -1.0);
assert_eq!(i.clamp(5.0), 2.0);
assert_eq!(i.clamp(0.25), 0.25);
assert!(i.clamp(f64::NAN).is_nan());
assert_eq!(Interval::REAL.clamp(1e300), 1e300);
assert!(Interval::REAL.contains(f64::INFINITY));
}
#[test]
fn lerp_hits_the_ends_exactly() {
let i = iv(0.1, 0.7);
assert_eq!(i.lerp(0.0), 0.1);
assert_eq!(i.lerp(1.0), 0.7);
assert!((i.midpoint() - 0.4).abs() <= f64::EPSILON);
assert_eq!(i.length(), 0.7 - 0.1);
assert!(i.is_bounded() && !Interval::REAL.is_bounded());
assert!(Interval::REAL.midpoint().is_nan());
}
#[test]
fn overlaps_intersection_hull() {
let a = iv(0.0, 1.0);
let b = iv(1.0, 2.0);
let c = iv(1.5, 3.0);
assert!(a.overlaps(&b) && b.overlaps(&a));
assert!(!a.overlaps(&c) && !c.overlaps(&a));
assert_eq!(a.intersection(&b), Some(iv(1.0, 1.0)));
assert_eq!(a.intersection(&c), None);
assert_eq!(b.intersection(&c), Some(iv(1.5, 2.0)));
assert_eq!(a.hull(&c), iv(0.0, 3.0));
assert_eq!(Interval::REAL.intersection(&c), Some(c));
assert_eq!(Interval::UNIT, a);
assert_eq!(Interval::TURN.hi(), core::f64::consts::TAU);
}
#[cfg(feature = "serde")]
#[test]
fn serde_round_trips_and_rejects_a_reversed_interval() {
let i = iv(-1.5, 2.25);
let text = serde_json::to_string(&i).unwrap();
assert_eq!(text, r#"{"lo":-1.5,"hi":2.25}"#);
let back: Interval = serde_json::from_str(&text).unwrap();
assert_eq!(back, i);
assert!(serde_json::from_str::<Interval>(r#"{"lo":2.0,"hi":1.0}"#).is_err());
}
}