use nalgebra::{DMatrix, DVector};
use crate::context::error::OxiflowError;
pub trait LinearSolver: Send + Sync {
fn solve(&self, a: &DMatrix<f64>, b: &DVector<f64>) -> Result<DVector<f64>, OxiflowError>;
}
#[derive(Debug, Default, Clone, Copy)]
pub struct NalgebraDenseSolver;
impl LinearSolver for NalgebraDenseSolver {
fn solve(&self, a: &DMatrix<f64>, b: &DVector<f64>) -> Result<DVector<f64>, OxiflowError> {
a.clone()
.lu()
.solve(b)
.ok_or_else(|| OxiflowError::PreconditionFailed {
context: "NalgebraDenseSolver",
message: "linear system is singular or near-singular (LU decomposition failed)"
.into(),
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn solves_identity_system() {
let a = DMatrix::<f64>::identity(3, 3);
let b = DVector::from_vec(vec![1.0, 2.0, 3.0]);
let x = NalgebraDenseSolver.solve(&a, &b).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 solves_diagonal_system() {
let a = DMatrix::from_diagonal(&DVector::from_vec(vec![2.0, 5.0]));
let b = DVector::from_vec(vec![4.0, 10.0]);
let x = NalgebraDenseSolver.solve(&a, &b).unwrap();
assert!((x[0] - 2.0).abs() < 1e-12);
assert!((x[1] - 2.0).abs() < 1e-12);
}
#[test]
fn singular_system_returns_error() {
let a = DMatrix::from_row_slice(2, 2, &[1.0, 1.0, 2.0, 2.0]);
let b = DVector::from_vec(vec![1.0, 2.0]);
let err = NalgebraDenseSolver.solve(&a, &b).unwrap_err();
assert!(matches!(err, OxiflowError::PreconditionFailed { .. }));
}
}