use crate::errors::QlResult;
use crate::math::array::Array;
use crate::methods::finitedifferences::operators::FdmLinearOpComposite;
use crate::methods::finitedifferences::utilities::FdmBoundaryConditionSet;
use crate::shared::SharedMut;
use crate::types::Time;
use crate::{fail, require};
use super::boundaryconditionschemehelper::BoundaryConditionSchemeHelper;
use super::scheme::Scheme;
pub struct ImplicitEulerScheme {
dt: Option<Time>,
map: SharedMut<dyn FdmLinearOpComposite>,
bc_set: BoundaryConditionSchemeHelper,
}
impl ImplicitEulerScheme {
pub fn new(map: SharedMut<dyn FdmLinearOpComposite>, bc_set: FdmBoundaryConditionSet) -> Self {
ImplicitEulerScheme {
dt: None,
map,
bc_set: BoundaryConditionSchemeHelper::new(bc_set),
}
}
}
impl Scheme for ImplicitEulerScheme {
fn set_step(&mut self, dt: Time) {
self.dt = Some(dt);
}
#[allow(clippy::neg_cmp_op_on_partial_ord)]
fn step(&mut self, a: &mut Array, t: Time) -> QlResult<()> {
let Some(dt) = self.dt else {
fail!("the timestep is not set: call set_step before stepping");
};
require!(t - dt > -1e-8, "a step towards negative time given");
let start = (t - dt).max(0.0);
{
let mut map = self.map.borrow_mut();
map.set_time(start, t)?;
self.bc_set.set_time(start);
self.bc_set.apply_before_solving(&mut *map, a);
let size = map.size();
if size != 1 {
fail!(
"implicit Euler over an operator splitting into {size} directions needs the \
iterative solvers deferred to #636"
);
}
*a = map.solve_splitting(0, a, -dt)?;
}
self.bc_set.apply_after_solving(a);
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use super::super::testops::{
GRID, assert_close, black_scholes_op, call_log, mesher, probe, scaled_composite,
};
use crate::shared::shared_mut;
use crate::types::Real;
const DT: Time = 0.1;
const T: Time = 0.25;
const COEFFICIENTS: [Real; 2] = [0.3, -0.45];
fn implicit_euler(
map: SharedMut<dyn FdmLinearOpComposite>,
bc_set: FdmBoundaryConditionSet,
) -> ImplicitEulerScheme {
let mut scheme = ImplicitEulerScheme::new(map, bc_set);
scheme.set_step(DT);
scheme
}
#[test]
fn a_step_is_the_implicit_solve_on_the_black_scholes_operator() {
let mesher = mesher();
let map: SharedMut<dyn FdmLinearOpComposite> = shared_mut(black_scholes_op(&mesher));
let mut scheme = implicit_euler(map, Vec::new());
let u = probe(GRID);
let mut a = u.clone();
scheme.step(&mut a, T).unwrap();
let mut oracle = black_scholes_op(&mesher);
oracle.set_time(T - DT, T).unwrap();
let expected = oracle.solve_splitting(0, &u, -DT).unwrap();
let moved = (0..u.size())
.map(|i| (expected[i] - u[i]).abs())
.fold(0.0, Real::max);
assert!(moved > 1e-6, "the solve leaves the values alone: {moved}");
assert_close(&a, &expected);
}
#[test]
fn a_step_matches_the_closed_form_on_a_diagonal_operator() {
let mut scheme = implicit_euler(scaled_composite(&COEFFICIENTS[..1]), Vec::new());
let u = probe(4);
let mut a = u.clone();
scheme.step(&mut a, T).unwrap();
assert_close(&a, &(&u / (1.0 - DT * COEFFICIENTS[0])));
}
#[test]
fn the_conditions_see_only_the_solving_calls_at_the_clamped_start() {
let raw = scaled_composite(&COEFFICIENTS[..1]);
let map: SharedMut<dyn FdmLinearOpComposite> = raw.clone();
let (log, bc_set) = call_log();
let mut scheme = implicit_euler(map, bc_set);
let t = DT - 5e-9;
scheme.step(&mut probe(4), t).unwrap();
assert_eq!(raw.borrow().last_set_time, Some((0.0, t)));
assert_eq!(
*log.borrow(),
vec![
"set_time:0".to_string(),
"before_solving".to_string(),
"after_solving".to_string(),
]
);
}
#[test]
fn a_multi_direction_operator_reports_the_deferred_iterative_solvers() {
let mut scheme = implicit_euler(scaled_composite(&COEFFICIENTS), Vec::new());
let error = scheme.step(&mut probe(4), T).unwrap_err();
assert!(error.message().contains("#636"), "{error}");
}
#[test]
fn stepping_before_the_timestep_is_set_fails() {
let mut scheme = ImplicitEulerScheme::new(scaled_composite(&COEFFICIENTS[..1]), Vec::new());
assert!(scheme.step(&mut probe(4), T).is_err());
}
#[test]
fn a_step_towards_negative_time_fails() {
let mut scheme = implicit_euler(scaled_composite(&COEFFICIENTS[..1]), Vec::new());
assert!(scheme.step(&mut probe(4), DT / 2.0).is_err());
}
}