use alloc::vec;
use alloc::vec::Vec;
use crate::math;
pub(crate) fn solve_cyclic_tridiagonal(
sub: &[f64],
diag: &[f64],
sup: &[f64],
corner_top_right: f64,
corner_bottom_left: f64,
rhs: &[f64],
) -> Option<Vec<f64>> {
let count = diag.len();
if count < 3 || sub.len() != count || sup.len() != count || rhs.len() != count {
return None;
}
let gamma = -diag.first()?;
if math::abs(gamma) < f64::EPSILON {
return None;
}
let ratio = corner_top_right / gamma;
let mut modified_diag = diag.to_vec();
*modified_diag.first_mut()? = diag.first()? - gamma;
*modified_diag.last_mut()? = diag.last()? - corner_bottom_left * ratio;
let mut correction = vec![0.0; count];
*correction.first_mut()? = gamma;
*correction.last_mut()? = corner_bottom_left;
let solved_rhs = solve_tridiagonal(sub, &modified_diag, sup, rhs)?;
let solved_correction = solve_tridiagonal(sub, &modified_diag, sup, &correction)?;
let numerator = solved_rhs.first()? + ratio * solved_rhs.last()?;
let denominator = 1.0 + solved_correction.first()? + ratio * solved_correction.last()?;
if math::abs(denominator) < f64::EPSILON {
return None;
}
let factor = numerator / denominator;
Some(
solved_rhs
.iter()
.zip(solved_correction.iter())
.map(|(&value, &adjustment)| value - factor * adjustment)
.collect(),
)
}
#[allow(clippy::indexing_slicing)]
pub(crate) fn solve_tridiagonal(
sub: &[f64],
diag: &[f64],
sup: &[f64],
rhs: &[f64],
) -> Option<Vec<f64>> {
let count = diag.len();
if count == 0 || sub.len() != count || sup.len() != count || rhs.len() != count {
return None;
}
let mut sweep = vec![0.0; count];
let mut solution = vec![0.0; count];
if math::abs(diag[0]) < f64::EPSILON {
return None;
}
sweep[0] = sup[0] / diag[0];
solution[0] = rhs[0] / diag[0];
for index in 1..count {
let pivot = diag[index] - sub[index] * sweep[index - 1];
if math::abs(pivot) < f64::EPSILON {
return None;
}
sweep[index] = sup[index] / pivot;
solution[index] = (rhs[index] - sub[index] * solution[index - 1]) / pivot;
}
for index in (0..count - 1).rev() {
solution[index] -= sweep[index] * solution[index + 1];
}
Some(solution)
}
#[allow(clippy::indexing_slicing)]
pub(crate) fn solve_dense(matrix: &mut [f64], rhs: &mut [f64], size: usize) -> Option<Vec<f64>> {
if size == 0 || matrix.len() != size * size || rhs.len() != size {
return None;
}
for column in 0..size {
let mut pivot_row = column;
let mut best = math::abs(matrix[column * size + column]);
for row in (column + 1)..size {
let candidate = math::abs(matrix[row * size + column]);
if candidate > best {
best = candidate;
pivot_row = row;
}
}
if best < 1e-12 {
return None;
}
if pivot_row != column {
for index in 0..size {
matrix.swap(column * size + index, pivot_row * size + index);
}
rhs.swap(column, pivot_row);
}
let pivot = matrix[column * size + column];
for row in (column + 1)..size {
let factor = matrix[row * size + column] / pivot;
if factor == 0.0 {
continue;
}
for index in column..size {
matrix[row * size + index] -= factor * matrix[column * size + index];
}
rhs[row] -= factor * rhs[column];
}
}
let mut solution = vec![0.0; size];
for row in (0..size).rev() {
let mut accumulator = rhs[row];
for column in (row + 1)..size {
accumulator -= matrix[row * size + column] * solution[column];
}
solution[row] = accumulator / matrix[row * size + row];
}
Some(solution)
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::float_cmp, clippy::indexing_slicing)]
mod tests {
use super::*;
#[test]
fn dense_solver_matches_a_hand_solution() {
let mut matrix = vec![2.0, 1.0, 1.0, 3.0];
let mut rhs = vec![5.0, 10.0];
let solution = solve_dense(&mut matrix, &mut rhs, 2).unwrap();
assert!((solution[0] - 1.0).abs() < 1e-12);
assert!((solution[1] - 3.0).abs() < 1e-12);
}
#[test]
fn dense_solver_pivots() {
let mut matrix = vec![0.0, 1.0, 1.0, 0.0];
let mut rhs = vec![2.0, 3.0];
let solution = solve_dense(&mut matrix, &mut rhs, 2).unwrap();
assert!((solution[0] - 3.0).abs() < 1e-12);
assert!((solution[1] - 2.0).abs() < 1e-12);
}
#[test]
fn dense_solver_rejects_singular_systems() {
let mut matrix = vec![1.0, 2.0, 2.0, 4.0];
let mut rhs = vec![1.0, 2.0];
assert!(solve_dense(&mut matrix, &mut rhs, 2).is_none());
assert!(solve_dense(&mut [], &mut [], 0).is_none());
}
#[test]
fn tridiagonal_solver_matches_a_hand_solution() {
let sub = [0.0, -1.0, -1.0];
let diag = [2.0, 2.0, 2.0];
let sup = [-1.0, -1.0, 0.0];
let rhs = [1.0, 0.0, 1.0];
let solution = solve_tridiagonal(&sub, &diag, &sup, &rhs).unwrap();
for value in solution {
assert!((value - 1.0).abs() < 1e-12);
}
}
#[test]
fn solvers_reject_mismatched_lengths() {
assert!(solve_tridiagonal(&[0.0], &[1.0, 2.0], &[0.0], &[1.0]).is_none());
assert!(solve_cyclic_tridiagonal(&[0.0], &[1.0], &[0.0], 1.0, 1.0, &[1.0]).is_none());
}
#[test]
fn cyclic_solver_matches_a_dense_solution() {
let sub = [0.0, 1.0, 1.0, 1.0];
let diag = [4.0, 4.0, 4.0, 4.0];
let sup = [1.0, 1.0, 1.0, 0.0];
let (corner_tr, corner_bl) = (1.0, 1.0);
let rhs = [1.0, 2.0, 3.0, 4.0];
let banded =
solve_cyclic_tridiagonal(&sub, &diag, &sup, corner_tr, corner_bl, &rhs).unwrap();
let mut dense = vec![0.0; 16];
for row in 0..4 {
dense[row * 4 + row] = diag[row];
if row > 0 {
dense[row * 4 + row - 1] = sub[row];
}
if row < 3 {
dense[row * 4 + row + 1] = sup[row];
}
}
dense[3] = corner_tr;
dense[12] = corner_bl;
let mut dense_rhs = rhs.to_vec();
let reference = solve_dense(&mut dense, &mut dense_rhs, 4).unwrap();
for (left, right) in banded.iter().zip(reference.iter()) {
assert!((left - right).abs() < 1e-10, "{left} vs {right}");
}
}
}