use super::MathError;
const PIVOT_TOLERANCE: f64 = 1e-14;
#[allow(clippy::many_single_char_names)]
pub fn thomas(a: &[f64], b: &[f64], c: &[f64], d: &[f64]) -> Result<Vec<f64>, MathError> {
let n = b.len();
if n == 0 || a.len() != n || c.len() != n || d.len() != n {
return Err(MathError::DimensionMismatch);
}
let mut c_prime = vec![0.0_f64; n];
let mut d_prime = vec![0.0_f64; n];
if b[0].abs() < PIVOT_TOLERANCE {
return Err(MathError::Singular);
}
c_prime[0] = c[0] / b[0];
d_prime[0] = d[0] / b[0];
for i in 1..n {
let denom = b[i] - a[i] * c_prime[i - 1];
if denom.abs() < PIVOT_TOLERANCE {
return Err(MathError::Singular);
}
if i < n - 1 {
c_prime[i] = c[i] / denom;
}
d_prime[i] = (d[i] - a[i] * d_prime[i - 1]) / denom;
}
let mut x = vec![0.0_f64; n];
x[n - 1] = d_prime[n - 1];
for i in (0..(n - 1)).rev() {
x[i] = d_prime[i] - c_prime[i] * x[i + 1];
}
Ok(x)
}
#[cfg(test)]
#[allow(clippy::many_single_char_names)]
mod tests {
use super::*;
use crate::math::linear_solve::solve;
#[test]
fn thomas_3x3_known() {
let a = vec![0.0, 1.0, 1.0];
let b = vec![2.0, 2.0, 2.0];
let c = vec![1.0, 1.0, 0.0];
let d = vec![4.0, 8.0, 8.0];
let x = thomas(&a, &b, &c, &d).unwrap();
assert!((x[0] - 1.0).abs() < 1e-12);
assert!((x[1] - 2.0).abs() < 1e-12);
assert!((x[2] - 3.0).abs() < 1e-12);
}
#[test]
fn thomas_4x4_diagonally_dominant() {
let a = vec![0.0, -1.0, -1.0, -1.0];
let b = vec![4.0, 4.0, 4.0, 4.0];
let c = vec![-1.0, -1.0, -1.0, 0.0];
let d = vec![5.0, 5.0, 10.0, 15.0];
let x = thomas(&a, &b, &c, &d).unwrap();
let r0 = 4.0 * x[0] - x[1] - 5.0;
let r1 = -x[0] + 4.0 * x[1] - x[2] - 5.0;
let r2 = -x[1] + 4.0 * x[2] - x[3] - 10.0;
let r3 = -x[2] + 4.0 * x[3] - 15.0;
for &r in &[r0, r1, r2, r3] {
assert!(r.abs() < 1e-12);
}
}
#[test]
fn thomas_n2_degenerate() {
let a = vec![0.0, 4.0];
let b = vec![2.0, 5.0];
let c = vec![3.0, 0.0];
let d = vec![13.0, 23.0];
let x = thomas(&a, &b, &c, &d).unwrap();
assert!((x[0] - 2.0).abs() < 1e-12);
assert!((x[1] - 3.0).abs() < 1e-12);
}
#[test]
fn thomas_n1_trivial() {
let a = vec![0.0];
let b = vec![5.0];
let c = vec![0.0];
let d = vec![10.0];
let x = thomas(&a, &b, &c, &d).unwrap();
assert!((x[0] - 2.0).abs() < 1e-15);
}
#[test]
fn thomas_dimension_mismatch() {
let a = vec![0.0, 1.0];
let b = vec![2.0, 2.0, 2.0];
let c = vec![1.0, 1.0, 0.0];
let d = vec![4.0, 8.0, 8.0];
let err = thomas(&a, &b, &c, &d).unwrap_err();
assert!(matches!(err, MathError::DimensionMismatch));
}
#[test]
fn thomas_empty_rejected() {
let empty: Vec<f64> = vec![];
let err = thomas(&empty, &empty, &empty, &empty).unwrap_err();
assert!(matches!(err, MathError::DimensionMismatch));
}
#[test]
fn thomas_zero_pivot_rejected() {
let a = vec![0.0, 1.0];
let b = vec![0.0, 2.0];
let c = vec![1.0, 0.0];
let d = vec![1.0, 1.0];
let err = thomas(&a, &b, &c, &d).unwrap_err();
assert!(matches!(err, MathError::Singular));
}
#[test]
fn thomas_singular_during_sweep() {
let a = vec![0.0, 1.0];
let b = vec![1.0, 2.0];
let c = vec![2.0, 0.0];
let d = vec![1.0, 1.0];
let err = thomas(&a, &b, &c, &d).unwrap_err();
assert!(matches!(err, MathError::Singular));
}
#[test]
fn thomas_cross_check_with_dense_solve_6x6() {
let sub = [-1.0_f64; 6];
let diag = [2.0_f64; 6];
let sup = [-1.0_f64; 6];
let a = sub.to_vec();
let b = diag.to_vec();
let c = sup.to_vec();
let d_rhs = vec![1.0, 0.0, 0.0, 0.0, 0.0, 1.0];
let x_thomas = thomas(&a, &b, &c, &d_rhs).unwrap();
let mut dense = vec![vec![0.0_f64; 6]; 6];
for i in 0..6 {
dense[i][i] = b[i];
if i + 1 < 6 {
dense[i][i + 1] = c[i];
dense[i + 1][i] = a[i + 1];
}
}
let mut rhs = d_rhs.clone();
solve(&mut dense, &mut rhs).unwrap();
for (t, g) in x_thomas.iter().zip(rhs.iter()) {
assert!((t - g).abs() < 1e-10);
}
}
}