#[inline(always)]
pub fn s_funct(psi: f64, alpha: f64) -> (f64, f64, f64, f64) {
const MAX_SERIES_TERMS: usize = 70;
const MAX_HALVING_STEPS: usize = 30;
const BETA_SERIES_THRESHOLD: f64 = 100.0;
let machine_epsilon = f64::EPSILON;
let series_term_convergence_tolerance = 100.0 * machine_epsilon;
let series_term_overflow_limit = 1.0 / machine_epsilon;
if psi == 0.0 {
return (1.0, 0.0, 0.0, 0.0);
}
let psi_squared = psi * psi;
let beta = alpha * psi_squared;
if beta.abs() < BETA_SERIES_THRESHOLD {
compute_stumpff_via_power_series(
psi,
psi_squared,
beta,
alpha,
series_term_convergence_tolerance,
series_term_overflow_limit,
MAX_SERIES_TERMS,
)
} else {
compute_stumpff_via_halving_and_duplication(
psi,
beta,
alpha,
series_term_convergence_tolerance,
series_term_overflow_limit,
BETA_SERIES_THRESHOLD,
MAX_HALVING_STEPS,
MAX_SERIES_TERMS,
)
}
}
#[allow(clippy::too_many_arguments)]
fn compute_stumpff_via_power_series(
psi: f64,
psi_squared: f64,
beta: f64,
alpha: f64,
convergence_tolerance: f64,
overflow_limit: f64,
max_terms: usize,
) -> (f64, f64, f64, f64) {
let mut s2 = 0.5 * psi_squared;
let mut series_term_s2 = s2;
let mut s3 = (s2 * psi) / 3.0;
let mut series_term_s3 = s3;
let mut denominator_s2_low = 3.0;
let mut denominator_s2_high = 4.0;
let mut denominator_s3_low = 4.0;
let mut denominator_s3_high = 5.0;
for _ in 1..=max_terms {
series_term_s2 *= beta / (denominator_s2_low * denominator_s2_high);
s2 += series_term_s2;
series_term_s3 *= beta / (denominator_s3_low * denominator_s3_high);
s3 += series_term_s3;
let term_s2_is_negligible = series_term_s2.abs() < convergence_tolerance;
let term_s3_is_negligible = series_term_s3.abs() < convergence_tolerance;
let term_s2_is_diverging = series_term_s2.abs() > overflow_limit;
let term_s3_is_diverging = series_term_s3.abs() > overflow_limit;
if (term_s2_is_negligible && term_s3_is_negligible)
|| term_s2_is_diverging
|| term_s3_is_diverging
{
break;
}
denominator_s2_low += 2.0;
denominator_s2_high += 2.0;
denominator_s3_low += 2.0;
denominator_s3_high += 2.0;
}
let s1 = psi + alpha * s3;
let s0 = 1.0 + alpha * s2;
(s0, s1, s2, s3)
}
#[allow(clippy::too_many_arguments)]
fn compute_stumpff_via_halving_and_duplication(
psi: f64,
beta: f64,
alpha: f64,
convergence_tolerance: f64,
overflow_limit: f64,
beta_threshold: f64,
max_halving_steps: usize,
max_terms: usize,
) -> (f64, f64, f64, f64) {
let (reduced_psi, reduced_beta, halving_count) =
reduce_psi_until_beta_is_small(psi, beta, beta_threshold, max_halving_steps);
let (mut s0, mut s1) = expand_s0_s1_series_at_reduced_psi(
reduced_psi,
reduced_beta,
convergence_tolerance,
overflow_limit,
max_terms,
);
for _ in 0..halving_count {
let cosine_like = s0;
let sine_like = s1;
s0 = 2.0 * cosine_like * cosine_like - 1.0;
s1 = 2.0 * cosine_like * sine_like;
}
let s3 = (s1 - psi) / alpha;
let s2 = (s0 - 1.0) / alpha;
(s0, s1, s2, s3)
}
fn reduce_psi_until_beta_is_small(
psi: f64,
beta: f64,
beta_threshold: f64,
max_halving_steps: usize,
) -> (f64, f64, usize) {
let mut reduced_psi = psi;
let mut reduced_beta = beta;
let mut halving_count = 0usize;
while reduced_beta.abs() >= beta_threshold && halving_count < max_halving_steps {
reduced_psi *= 0.5;
reduced_beta *= 0.25; halving_count += 1;
}
(reduced_psi, reduced_beta, halving_count)
}
fn expand_s0_s1_series_at_reduced_psi(
reduced_psi: f64,
reduced_beta: f64,
convergence_tolerance: f64,
overflow_limit: f64,
max_terms: usize,
) -> (f64, f64) {
let mut s0 = 1.0;
let mut s1 = reduced_psi;
let mut series_term_s0 = 1.0;
let mut series_term_s1 = reduced_psi;
for term_index in 1..=max_terms {
series_term_s0 *= reduced_beta / ((2 * term_index - 1) as f64 * (2 * term_index) as f64);
s0 += series_term_s0;
if series_term_s0.abs() < convergence_tolerance || series_term_s0.abs() > overflow_limit {
break;
}
}
for term_index in 1..=max_terms {
series_term_s1 *= reduced_beta / ((2 * term_index) as f64 * (2 * term_index + 1) as f64);
s1 += series_term_s1;
if series_term_s1.abs() < convergence_tolerance || series_term_s1.abs() > overflow_limit {
break;
}
}
(s0, s1)
}
#[cfg(test)]
mod tests_s_funct {
use approx::assert_relative_eq;
use super::s_funct;
fn check_invariants(psi: f64, alpha: f64, s0: f64, s1: f64, s2: f64, s3: f64) {
let tol = 1e-12;
assert!(
(s0 - (1.0 + alpha * s2)).abs() < tol,
"Invariant s0 = 1 + α*s2 violated: {} vs {}",
s0,
1.0 + alpha * s2
);
assert!(
(s1 - (psi + alpha * s3)).abs() < tol,
"Invariant s1 = ψ + α*s3 violated: {} vs {}",
s1,
psi + alpha * s3
);
}
#[test]
fn test_small_beta() {
let psi = 0.01;
let alpha = 0.1;
let (s0, s1, s2, s3) = s_funct(psi, alpha);
assert!(s0 > 0.0);
assert!(s1 > 0.0);
check_invariants(psi, alpha, s0, s1, s2, s3);
}
#[test]
fn test_large_beta() {
let psi = 10.0;
let alpha = 5.0;
let (s0, s1, s2, s3) = s_funct(psi, alpha);
assert!(s0.is_finite() && s1.is_finite() && s2.is_finite() && s3.is_finite());
let rel_tol = 1e-7;
assert_relative_eq!(s0, 1.0 + alpha * s2, max_relative = rel_tol);
assert_relative_eq!(s1, psi + alpha * s3, max_relative = rel_tol);
}
#[test]
fn test_zero_alpha() {
let psi = 2.0;
let alpha = 0.0;
let (s0, s1, s2, s3) = s_funct(psi, alpha);
assert!((s0 - 1.0).abs() < 1e-14);
assert!((s1 - psi).abs() < 1e-14);
assert!((s2 - psi.powi(2) / 2.0).abs() < 1e-14);
assert!((s3 - psi.powi(3) / 6.0).abs() < 1e-14);
}
#[test]
fn test_zero_psi() {
let psi = 0.0;
let alpha = 2.0;
let (s0, s1, s2, s3) = s_funct(psi, alpha);
assert!((s0 - 1.0).abs() < 1e-14);
assert!((s1 - 0.0).abs() < 1e-14);
assert!((s2 - 0.0).abs() < 1e-14);
assert!((s3 - 0.0).abs() < 1e-14);
}
#[test]
fn test_symmetry_negative_psi() {
let psi = 1.0;
let alpha = 0.5;
let (s0_pos, s1_pos, s2_pos, s3_pos) = s_funct(psi, alpha);
let (s0_neg, s1_neg, s2_neg, s3_neg) = s_funct(-psi, alpha);
let tol = 1e-12;
assert!((s0_pos - s0_neg).abs() < tol);
assert!((s2_pos - s2_neg).abs() < tol);
assert!((s1_pos + s1_neg).abs() < tol);
assert!((s3_pos + s3_neg).abs() < tol);
}
#[test]
fn test_consistency_large_vs_small() {
let psi = 2.5;
let alpha = 1.0;
let (s0, s1, s2, s3) = s_funct(psi, alpha);
check_invariants(psi, alpha, s0, s1, s2, s3);
}
#[test]
fn test_s_funct_real_data() {
let psi = -15.279808141051223;
let alpha = -1.6298946008705195e-4;
let (s0, s1, s2, s3) = s_funct(psi, alpha);
assert_eq!(s0, 0.9810334785583247);
assert_eq!(s1, -15.183083836892674);
assert_eq!(s2, 116.3665517484714);
assert_eq!(s3, -593.4390119881925);
}
}