use super::methods::backward_euler::BackwardEulerSolver;
use super::sparse::FaerSparseSolver;
use super::Solver;
use crate::context::error::OxiflowError;
#[non_exhaustive]
#[derive(Debug, Clone)]
#[cfg_attr(feature = "serde", derive(serde::Deserialize))]
pub enum IntegratorSpec {
BackwardEuler {
#[cfg_attr(feature = "serde", serde(default = "default_sparse_threshold"))]
sparse_threshold: usize,
#[cfg_attr(feature = "serde", serde(default))]
jacobian_bandwidth: Option<usize>,
},
}
#[cfg(feature = "serde")]
fn default_sparse_threshold() -> usize {
100
}
impl TryFrom<IntegratorSpec> for Box<dyn Solver> {
type Error = OxiflowError;
fn try_from(spec: IntegratorSpec) -> Result<Self, Self::Error> {
match spec {
IntegratorSpec::BackwardEuler {
sparse_threshold,
jacobian_bandwidth,
} => {
let mut solver = BackwardEulerSolver::new().with_sparse_threshold(sparse_threshold);
if let Some(bandwidth) = jacobian_bandwidth {
solver = solver
.with_jacobian_bandwidth(bandwidth)
.with_sparse_solver(Box::new(FaerSparseSolver));
}
Ok(Box::new(solver))
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn dense_path_when_no_jacobian_bandwidth() {
let spec = IntegratorSpec::BackwardEuler {
sparse_threshold: 100,
jacobian_bandwidth: None,
};
let solver: Box<dyn Solver> = spec.try_into().unwrap();
let _ = solver;
}
#[test]
fn sparse_path_when_jacobian_bandwidth_given() {
let spec = IntegratorSpec::BackwardEuler {
sparse_threshold: 50,
jacobian_bandwidth: Some(3),
};
let solver: Box<dyn Solver> = spec.try_into().unwrap();
let _ = solver;
}
#[cfg(feature = "serde")]
#[test]
fn deserialises_with_defaults() {
let json = r#"{ "BackwardEuler": {} }"#;
let spec: IntegratorSpec = serde_json::from_str(json).unwrap();
match spec {
IntegratorSpec::BackwardEuler {
sparse_threshold,
jacobian_bandwidth,
} => {
assert_eq!(sparse_threshold, default_sparse_threshold());
assert_eq!(jacobian_bandwidth, None);
}
}
}
#[cfg(feature = "serde")]
#[test]
fn deserialises_explicit_fields() {
let json = r#"{ "BackwardEuler": { "sparse_threshold": 42, "jacobian_bandwidth": 2 } }"#;
let spec: IntegratorSpec = serde_json::from_str(json).unwrap();
match spec {
IntegratorSpec::BackwardEuler {
sparse_threshold,
jacobian_bandwidth,
} => {
assert_eq!(sparse_threshold, 42);
assert_eq!(jacobian_bandwidth, Some(2));
}
}
}
}