use crate::core::curves::Compounding;
use crate::core::data_models::EquityOptionData;
use crate::core::errors::RustyQLibError;
use crate::core::fd_solvers::brennan_schwartz;
pub use crate::core::fd_solvers::thomas_algorithm;
use crate::core::trade::PutOrCall;
use crate::core::utils::ContractStyle;
use crate::equity::barrier::{BarrierDirection, KnockType};
use crate::equity::local_vol::LocalVol;
use crate::equity::utils::Model;
use crate::equity::utils::Payoff;
use crate::equity::vanilla_option::{BarrierPayoff, EquityOption};
const RANNACHER_STEPS: usize = 4;
const CELL_AVG_POINTS: usize = 16;
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct FdConfig {
pub spot_steps: usize,
pub time_steps: usize,
pub grid_stdevs: f64,
}
impl Default for FdConfig {
fn default() -> Self {
FdConfig { spot_steps: 400, time_steps: 400, grid_stdevs: 5.0 }
}
}
impl FdConfig {
pub fn from_data(data: &EquityOptionData) -> Self {
let defaults = FdConfig::default();
FdConfig {
spot_steps: data.fd_spot_steps.unwrap_or(defaults.spot_steps).max(10),
time_steps: data.fd_time_steps.unwrap_or(defaults.time_steps).max(10),
grid_stdevs: defaults.grid_stdevs,
}
}
pub fn validate(&self) -> Result<(), RustyQLibError> {
if self.spot_steps < 3 || self.time_steps < 1 {
return Err(RustyQLibError::invalid_input(
"fd_grid",
format!(
"the FD grid needs at least 3 spot steps and 1 time step, got {} x {}",
self.spot_steps, self.time_steps
),
));
}
Ok(())
}
}
#[derive(Debug, Clone, Copy)]
pub struct FdSolution {
pub npv: f64,
pub delta: f64,
pub gamma: f64,
pub theta: f64,
}
impl FdSolution {
fn zero() -> Self {
FdSolution { npv: 0.0, delta: 0.0, gamma: 0.0, theta: 0.0 }
}
fn minus(self, other: FdSolution) -> Self {
FdSolution {
npv: self.npv - other.npv,
delta: self.delta - other.delta,
gamma: self.gamma - other.gamma,
theta: self.theta - other.theta,
}
}
}
pub fn npv(option: &EquityOption) -> f64 {
solution(option).npv
}
pub fn delta(option: &EquityOption) -> f64 {
solution(option).delta
}
pub fn gamma(option: &EquityOption) -> f64 {
solution(option).gamma
}
pub fn theta(option: &EquityOption) -> f64 {
solution(option).theta
}
pub fn vega(option: &EquityOption) -> f64 {
let h = 1e-3;
(solve_dispatch(option, h, 0.0, 0.0).npv - solve_dispatch(option, -h, 0.0, 0.0).npv) / (2.0 * h)
}
pub fn rho(option: &EquityOption) -> f64 {
let h = 1e-4;
(solve_dispatch(option, 0.0, h, 0.0).npv - solve_dispatch(option, 0.0, -h, 0.0).npv) / (2.0 * h)
}
pub fn vanna(option: &EquityOption) -> f64 {
let h = 1e-3;
(solve_dispatch(option, h, 0.0, 0.0).delta - solve_dispatch(option, -h, 0.0, 0.0).delta)
/ (2.0 * h)
}
pub fn charm(option: &EquityOption) -> f64 {
let h = option.market.spot.value() * 1e-3;
(solve_dispatch(option, 0.0, 0.0, h).theta - solve_dispatch(option, 0.0, 0.0, -h).theta)
/ (2.0 * h)
}
pub fn zomma(option: &EquityOption) -> f64 {
let h = 1e-3;
(solve_dispatch(option, h, 0.0, 0.0).gamma - solve_dispatch(option, -h, 0.0, 0.0).gamma)
/ (2.0 * h)
}
pub fn volga(option: &EquityOption) -> f64 {
let h = 1e-2;
(solve_dispatch(option, h, 0.0, 0.0).npv - 2.0 * solve_dispatch(option, 0.0, 0.0, 0.0).npv
+ solve_dispatch(option, -h, 0.0, 0.0).npv)
/ (h * h)
}
pub fn solution(option: &EquityOption) -> FdSolution {
solve_dispatch(option, 0.0, 0.0, 0.0)
}
pub fn pricing_result(option: &EquityOption) -> crate::core::results::PricingResult {
use crate::core::results::{Greeks, PricingResult};
let base = solution(option);
let hv = 1e-3;
let vol_up = solve_dispatch(option, hv, 0.0, 0.0);
let vol_down = solve_dispatch(option, -hv, 0.0, 0.0);
let hr = 1e-4;
let rho = (solve_dispatch(option, 0.0, hr, 0.0).npv
- solve_dispatch(option, 0.0, -hr, 0.0).npv)
/ (2.0 * hr);
let hs = option.market.spot.value() * 1e-3;
let charm = (solve_dispatch(option, 0.0, 0.0, hs).theta
- solve_dispatch(option, 0.0, 0.0, -hs).theta)
/ (2.0 * hs);
let gamma_p = if base.delta == 0.0 {
f64::NAN
} else {
option.market.spot.value() * base.gamma / base.delta
};
PricingResult {
pv: base.npv,
greeks: Greeks {
delta: base.delta,
gamma: base.gamma,
vega: (vol_up.npv - vol_down.npv) / (2.0 * hv),
theta: base.theta,
rho,
vanna: (vol_up.delta - vol_down.delta) / (2.0 * hv),
charm,
gamma_p,
zomma: (vol_up.gamma - vol_down.gamma) / (2.0 * hv),
},
std_err: None,
}
}
pub(crate) fn npv_with(
option: &EquityOption,
d_spot: f64,
d_vol: f64,
d_rate: f64,
d_time: f64,
) -> f64 {
let sol = solve_dispatch(option, d_vol, d_rate, d_spot);
sol.npv + sol.theta * d_time
}
fn solve_dispatch(
option: &EquityOption,
sigma_bump: f64,
r_bump: f64,
spot_bump: f64,
) -> FdSolution {
let t = option.time_to_maturity();
assert!(t >= 0.0, "Option is expired or negative time");
let s0 = option.market.spot.value() + spot_bump;
assert!(s0 > 0.0, "underlying price must be positive");
if t == 0.0 {
let mut sol = FdSolution::zero();
sol.npv = option.payoff.payoff(s0, option.base.strike_price);
return sol;
}
if option.model.is_heston() {
return super::heston_adi::solve(option, sigma_bump, r_bump, spot_bump);
}
if let Some(barrier) = option.payoff.as_any().downcast_ref::<BarrierPayoff>() {
assert!(
barrier.barrier2.is_none() && barrier.rebate == 0.0,
"double barriers and rebates are not supported on the FD engine; use the Analytical or MonteCarlo engine"
);
let down = barrier.direction == BarrierDirection::Down;
let knocked = if down { s0 <= barrier.barrier } else { s0 >= barrier.barrier };
return match barrier.knock {
KnockType::Out => {
if knocked {
FdSolution::zero()
} else {
solve(option, sigma_bump, r_bump, spot_bump, Some(barrier))
}
}
KnockType::In => {
let vanilla = solve(option, sigma_bump, r_bump, spot_bump, None);
if knocked {
vanilla
} else {
vanilla.minus(solve(option, sigma_bump, r_bump, spot_bump, Some(barrier)))
}
}
};
}
solve(option, sigma_bump, r_bump, spot_bump, None)
}
enum FdVol<'a> {
Const(f64),
Local(LocalVol<'a>),
}
impl FdVol<'_> {
fn vol(&self, s: f64, calendar_t: f64) -> f64 {
match self {
FdVol::Const(v) => *v,
FdVol::Local(lv) => lv.vol(s, calendar_t),
}
}
}
fn solve(
option: &EquityOption,
sigma_bump: f64,
r_bump: f64,
spot_bump: f64,
knock_out: Option<&BarrierPayoff>,
) -> FdSolution {
let cfg = option.fd_cfg();
let payoff = option.payoff.as_ref();
let strike = option.base.strike_price;
let s0 = option.market.spot.value() + spot_bump;
let q = option.carry_yield();
let t = option.time_to_maturity();
let sigma_ref = option.volatility() + sigma_bump;
assert!(sigma_ref > 0.0, "volatility must be positive");
let american = matches!(payoff.exercise_style(), ContractStyle::American);
let bermudan_backward: Option<Vec<bool>> = match payoff.exercise_style() {
ContractStyle::Bermudan(times) => {
let steps_total = cfg.time_steps;
let mut mask = vec![false; steps_total];
for g in crate::core::utils::times_to_grid_steps(times, t, steps_total) {
if g < steps_total {
mask[steps_total - g - 1] = true;
}
}
Some(mask)
}
_ => None,
};
let put = matches!(payoff.put_or_call(), PutOrCall::Put);
let vol_field = match option.model {
Model::Gbm => FdVol::Const(sigma_ref),
Model::LocalVol => FdVol::Local(LocalVol::new(
&option.market.vol_surface,
&option.market.discount_curve,
option.market.spot.value(),
q,
sigma_bump,
)),
Model::Heston(_) => unreachable!("Heston is dispatched to heston_adi::solve"),
};
let x0 = s0.ln();
let r_flat = option.risk_free_rate() + r_bump;
let drift_width = ((r_flat - q - 0.5 * sigma_ref * sigma_ref) * t).abs();
let half_width =
cfg.grid_stdevs * sigma_ref * t.sqrt() + drift_width + (strike / s0).ln().abs().max(1e-2);
let (x_min, x_max, barrier_low, barrier_high) = match knock_out {
Some(b) if b.direction == BarrierDirection::Down => {
(b.barrier.ln(), x0 + half_width, true, false)
}
Some(b) => (x0 - half_width, b.barrier.ln(), false, true),
None => (x0 - half_width, x0 + half_width, false, false),
};
let n = cfg.spot_steps;
let dx = (x_max - x_min) / n as f64;
let x_at = |i: usize| x_min + i as f64 * dx;
let s_grid: Vec<f64> = (0..=n).map(|i| x_at(i).exp()).collect();
let exercise: Vec<f64> = s_grid.iter().map(|&s| payoff.payoff(s, strike)).collect();
let steps = cfg.time_steps;
let dt = t / steps as f64;
let curve = &option.market.discount_curve;
let step_rates: Vec<f64> = (0..steps)
.map(|k| {
let t2 = t - k as f64 * dt;
let t1 = t - (k + 1) as f64 * dt;
let fwd = if t1 <= 0.0 {
curve.zero_rate_with(t2.max(1e-8), Compounding::Continuous)
} else {
curve
.forward_rate_with(t1, t2, Compounding::Continuous)
.unwrap_or_else(|_| curve.zero_rate_with(t2, Compounding::Continuous))
};
fwd + r_bump
})
.collect();
let cash_divs: Vec<(f64, f64)> = option
.market
.cash_dividends
.iter()
.filter_map(|(date, amount)| {
let td = (*date - option.market.valuation_date).num_days() as f64 / 365.0;
(td > 0.0 && td <= t).then_some((td, *amount))
})
.collect();
let mut v: Vec<f64> = (0..=n)
.map(|i| cell_average_payoff(payoff, strike, x_at(i), dx))
.collect();
if barrier_low {
v[0] = 0.0;
}
if barrier_high {
v[n] = 0.0;
}
let m = n - 1; let mut sub = vec![0.0; m - 1];
let mut dia = vec![0.0; m];
let mut sup = vec![0.0; m - 1];
let mut rhs = vec![0.0; m];
let mut lower = vec![0.0; n + 1];
let mut diag = vec![0.0; n + 1];
let mut upper = vec![0.0; n + 1];
let mut cum_df = 1.0;
let mut cum_growth = 1.0;
let mut theta_layer_value = 0.0;
for step in 0..steps {
let exercise_now = american
|| bermudan_backward.as_ref().map_or(false, |m| m.get(step).copied().unwrap_or(false));
let theta_w = if step < RANNACHER_STEPS { 1.0 } else { 0.5 };
let r_step = step_rates[step];
let calendar_mid = (t - (step as f64 + 0.5) * dt).max(0.0);
cum_df *= (-r_step * dt).exp();
cum_growth *= ((r_step - q) * dt).exp();
for i in 0..=n {
let sigma = vol_field.vol(s_grid[i], calendar_mid);
let s2 = 0.5 * sigma * sigma;
let mu = r_step - q - s2;
lower[i] = s2 / (dx * dx) - mu / (2.0 * dx);
diag[i] = -2.0 * s2 / (dx * dx) - r_step;
upper[i] = s2 / (dx * dx) + mu / (2.0 * dx);
}
let boundary = |i: usize, is_barrier: bool| -> f64 {
if is_barrier {
return 0.0;
}
let mut val = cum_df * payoff.payoff(s_grid[i] * cum_growth, strike);
if exercise_now {
val = val.max(exercise[i]);
}
val
};
let v_low = boundary(0, barrier_low);
let v_high = boundary(n, barrier_high);
for i in 1..n {
let av = lower[i] * v[i - 1] + diag[i] * v[i] + upper[i] * v[i + 1];
rhs[i - 1] = v[i] + (1.0 - theta_w) * dt * av;
}
rhs[0] += theta_w * dt * lower[1] * v_low;
rhs[m - 1] += theta_w * dt * upper[n - 1] * v_high;
for i in 1..n {
dia[i - 1] = 1.0 - theta_w * dt * diag[i];
}
for i in 1..n - 1 {
sub[i - 1] = -theta_w * dt * lower[i + 1];
sup[i - 1] = -theta_w * dt * upper[i];
}
let interior = if exercise_now {
brennan_schwartz(&sub, &dia, &sup, &rhs, &exercise[1..n], put)
} else {
thomas_algorithm(&sub, &dia, &sup, &rhs)
};
v[0] = v_low;
v[n] = v_high;
v[1..n].copy_from_slice(&interior);
if !cash_divs.is_empty() {
let cal_old = t - step as f64 * dt;
let cal_new = t - (step + 1) as f64 * dt;
let crossing: f64 = cash_divs
.iter()
.filter(|(td, _)| *td < cal_old && *td >= cal_new)
.map(|(_, amount)| *amount)
.sum();
if crossing > 0.0 {
let shifted: Vec<f64> = (0..=n)
.map(|i| {
let s_target = s_grid[i] - crossing;
if s_target <= s_grid[0] {
v[0]
} else {
let x_target = s_target.ln();
let j =
(((x_target - x_min) / dx).floor() as usize).min(n - 1);
let w = ((x_target - x_at(j)) / dx).clamp(0.0, 1.0);
v[j] * (1.0 - w) + v[j + 1] * w
}
})
.collect();
v = shifted;
if exercise_now {
for i in 0..=n {
if v[i] < exercise[i] {
v[i] = exercise[i];
}
}
}
}
}
if step + 1 == steps.saturating_sub(1) {
theta_layer_value = read_grid(&v, x_min, dx, x0).0;
}
}
let (npv, delta_x, gamma_x) = read_grid(&v, x_min, dx, x0);
let delta = delta_x / s0;
let gamma = (gamma_x - delta_x) / (s0 * s0);
let theta = if steps >= 2 { (theta_layer_value - npv) / dt } else { 0.0 };
FdSolution { npv, delta, gamma, theta }
}
fn read_grid(v: &[f64], x_min: f64, dx: f64, x0: f64) -> (f64, f64, f64) {
let n = v.len() - 1;
let i = (((x0 - x_min) / dx).round() as usize).clamp(1, n - 1);
let e = x0 - (x_min + i as f64 * dx);
let b = (v[i + 1] - v[i - 1]) / (2.0 * dx);
let c = (v[i + 1] - 2.0 * v[i] + v[i - 1]) / (2.0 * dx * dx);
(v[i] + b * e + c * e * e, b + 2.0 * c * e, 2.0 * c)
}
fn cell_average_payoff(payoff: &dyn Payoff, strike: f64, x: f64, dx: f64) -> f64 {
let k = CELL_AVG_POINTS;
let mut sum = 0.0;
for j in 0..k {
let xi = x - 0.5 * dx + (j as f64 + 0.5) * dx / k as f64;
sum += payoff.payoff(xi.exp(), strike);
}
sum / k as f64
}