use super::number::Number;
type Error = Box<dyn std::error::Error>;
type Result<T> = std::result::Result<T, Error>;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Range<T: Number>(T, T);
impl<T: Number> Range<T> {
pub fn new(start: T, end: T) -> Result<Self> {
if start.is_nan() || end.is_nan() {
return Err("NaN is not allowed in range".into());
}
if start > end {
return Err("Invalid range (negative size)".into());
}
Ok(Range(start, end))
}
pub fn len(&self) -> T {
self.1 - self.0
}
pub fn contains(&self, value: T) -> bool {
value >= self.0 && value < self.1
}
pub fn try_merge(&self, other: &Self) -> Result<Self> {
if self.1 < other.0 || other.1 < self.0 {
return Err("Disjoint ranges cannot be merged".into());
}
Ok(Range(self.0.min(other.0), self.1.max(other.1)))
}
}
impl<T: Number> From<(T, T)> for Range<T> {
fn from(tuple: (T, T)) -> Self {
Range::new(tuple.0, tuple.1).unwrap()
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct RangeSet<T: Number> {
ranges: Vec<Range<T>>,
}
impl<T: Number> RangeSet<T> {
pub fn new() -> Self {
RangeSet { ranges: Vec::new() }
}
fn binary_search_by_first(&self, value: T) -> std::result::Result<usize, usize> {
self.ranges
.binary_search_by(|r| r.0.partial_cmp(&value).unwrap())
}
pub fn add_range<R: Into<Range<T>>>(&mut self, range: R) {
let range = range.into();
if range.len() == T::zero() {
return;
}
if self.ranges.is_empty() {
self.ranges.push(range);
return;
}
let start_pos = match self.binary_search_by_first(range.0) {
Ok(pos) => pos,
Err(0) => 0,
Err(pos) => {
if range.0 <= self.ranges[pos - 1].1 {
pos - 1
} else {
pos
}
}
};
let end_pos = match self.binary_search_by_first(range.1){
Ok(pos) => pos + 1,
Err(pos) => pos,
};
if start_pos == end_pos {
self.ranges.insert(start_pos, range);
} else {
let new_start = self.ranges[start_pos].0.min(range.0);
let new_end = self.ranges[end_pos - 1].1.max(range.1);
self.ranges[start_pos].0 = new_start;
self.ranges[start_pos].1 = new_end;
self.ranges.drain(start_pos + 1..end_pos);
}
}
pub fn contains(&self, value: T) -> bool {
match self.binary_search_by_first(value) {
Ok(_) => true,
Err(0) => false,
Err(pos) => self.ranges[pos - 1].contains(value),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_float_range_new() {
assert!(Range::new(1.0f64, 2.0).is_ok());
assert!(Range::new(2.0f64, 1.0).is_err());
assert!(Range::new(f64::NAN, 2.0).is_err());
}
#[test]
fn test_float_range_contains() {
let range = Range::new(1.0f64, 3.0).unwrap();
assert!(range.contains(1.0));
assert!(range.contains(2.0));
assert!(!range.contains(3.0));
}
#[test]
fn test_float_range_try_merge() {
let range1 = Range::new(1.0f64, 3.0).unwrap();
let range2 = Range::new(2.0f64, 4.0).unwrap();
let merged = range1.try_merge(&range2).unwrap();
assert_eq!(merged, Range::new(1.0, 4.0).unwrap());
let range3 = Range::new(4.0f64, 5.0).unwrap();
assert!(range1.try_merge(&range3).is_err());
assert!(range2.try_merge(&range3).is_ok());
}
#[test]
fn test_float_range_set_add_and_contains() {
let mut range_set = RangeSet::new();
range_set.add_range(Range::new(1.0f64, 3.0).unwrap());
range_set.add_range(Range::new(4.0f64, 7.0).unwrap());
range_set.add_range(Range::new(8.0f64, 10.0).unwrap());
let in_set = vec![1.0, 1.5, 2.0, 2.9, 4.0, 6.0, 8.0, 9.0];
for &value in &in_set {
assert!(range_set.contains(value));
}
let not_in_set = vec![-1.0, 0.0, 3.0, 3.5, 7.0, 7.5, 10.0, 11.0];
for &value in ¬_in_set {
assert!(!range_set.contains(value));
}
}
}