use nalgebra::ComplexField;
use num_traits::{Float, FromPrimitive};
use std::ops::Range;
#[derive(Debug, thiserror::Error)]
pub(crate) enum RangeSplitError<C: ComplexField> {
#[error("segment too small: minimum {minimum_segment_width:?} around {center:?}")]
SegmentTooSmall {
minimum_segment_width: C::RealField,
center: C,
},
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum SingularityHandling {
Error,
RecursiveSplit {
max_depth: usize,
},
}
impl Default for SingularityHandling {
fn default() -> Self {
Self::RecursiveSplit { max_depth: 32 }
}
}
impl SingularityHandling {
pub(super) fn split_range<I, F>(
&self,
range: Range<I>,
singularity: I,
minimum_segment_width: F,
) -> Result<[Range<I>; 2], RangeSplitError<I>>
where
I: ComplexField<RealField = F> + Copy,
F: Float + FromPrimitive,
{
let width = (range.end - range.start).modulus();
if width <= minimum_segment_width {
return Err(RangeSplitError::SegmentTooSmall {
minimum_segment_width,
center: singularity,
});
}
let two = I::from_real(F::one() + F::one());
let left_width = (singularity - range.start).modulus();
let right_width = (range.end - singularity).modulus();
if left_width <= minimum_segment_width || right_width <= minimum_segment_width {
let midpoint = (range.start + range.end) / two;
return Ok([range.start..midpoint, midpoint..range.end]);
}
Ok([range.start..singularity, singularity..range.end])
}
}
#[cfg(test)]
mod test {
use super::*;
const TOL: f64 = 1e-10;
fn assert_close(a: f64, b: f64) {
assert!(
(a - b).abs() < TOL,
"expected {b}, got {a}, diff = {}",
(a - b).abs()
);
}
#[test]
fn split_range_splits_at_interior_singularity() {
let policy = SingularityHandling::RecursiveSplit { max_depth: 8 };
let [left, right] = policy.split_range(0.0..10.0, 4.0, 1e-12).unwrap();
assert_close(left.start, 0.0);
assert_close(left.end, 4.0);
assert_close(right.start, 4.0);
assert_close(right.end, 10.0);
}
#[test]
fn split_range_falls_back_to_midpoint_when_singularity_is_too_close_to_edge() {
let policy = SingularityHandling::RecursiveSplit { max_depth: 8 };
let [left, right] = policy.split_range(0.0..10.0, 1e-14, 1e-12).unwrap();
assert_close(left.start, 0.0);
assert_close(left.end, 5.0);
assert_close(right.start, 5.0);
assert_close(right.end, 10.0);
}
#[test]
fn split_range_errors_when_range_is_too_small() {
let policy = SingularityHandling::RecursiveSplit { max_depth: 8 };
let result = policy.split_range(0.0..1e-14, 5e-15, 1e-12);
assert!(matches!(
result,
Err(RangeSplitError::SegmentTooSmall { .. })
));
}
}