use crate::derivatives::black_scholes::{
bsm_cross_greeks, bsm_greeks, bsm_price, BsmCrossGreeks, BsmGreeks,
};
use crate::derivatives::implied_vol::bsm_implied_vol;
use crate::derivatives::types::{validate_bsm_params, BsmParams, OptionType};
use crate::util::error::{require_finite, FinanceResult};
#[derive(Clone, Debug, PartialEq)]
pub struct BsmState {
params: BsmParams,
option_type: OptionType,
}
impl BsmState {
pub fn new(params: BsmParams, option_type: OptionType) -> FinanceResult<Self> {
validate_bsm_params(params)?;
Ok(Self {
params,
option_type,
})
}
pub fn params(&self) -> BsmParams {
self.params
}
pub fn option_type(&self) -> OptionType {
self.option_type
}
pub fn set_spot(&mut self, spot: f64) -> FinanceResult<()> {
require_finite("spot", spot)?;
let mut p = self.params;
p.spot = spot;
validate_bsm_params(p)?;
self.params = p;
Ok(())
}
pub fn set_vol(&mut self, vol: f64) -> FinanceResult<()> {
require_finite("vol", vol)?;
let mut p = self.params;
p.vol = vol;
validate_bsm_params(p)?;
self.params = p;
Ok(())
}
pub fn set_time_years(&mut self, time_years: f64) -> FinanceResult<()> {
require_finite("time_years", time_years)?;
let mut p = self.params;
p.time_years = time_years;
validate_bsm_params(p)?;
self.params = p;
Ok(())
}
pub fn set_strike(&mut self, strike: f64) -> FinanceResult<()> {
require_finite("strike", strike)?;
let mut p = self.params;
p.strike = strike;
validate_bsm_params(p)?;
self.params = p;
Ok(())
}
pub fn set_rate(&mut self, rate: f64) -> FinanceResult<()> {
require_finite("rate", rate)?;
let mut p = self.params;
p.rate = rate;
validate_bsm_params(p)?;
self.params = p;
Ok(())
}
pub fn set_dividend_yield(&mut self, dividend_yield: f64) -> FinanceResult<()> {
require_finite("dividend_yield", dividend_yield)?;
let mut p = self.params;
p.dividend_yield = dividend_yield;
validate_bsm_params(p)?;
self.params = p;
Ok(())
}
pub fn set_vol_from_price(&mut self, market_price: f64) -> FinanceResult<f64> {
let iv = bsm_implied_vol(self.params, self.option_type, market_price)?;
self.set_vol(iv)?;
Ok(iv)
}
pub fn price(&self) -> FinanceResult<f64> {
bsm_price(self.params, self.option_type)
}
pub fn greeks(&self) -> FinanceResult<BsmGreeks> {
bsm_greeks(self.params, self.option_type)
}
pub fn cross_greeks(&self) -> FinanceResult<BsmCrossGreeks> {
bsm_cross_greeks(self.params, self.option_type)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn set_spot_moves_delta() {
let p = BsmParams::atm_one_year(100.0, 0.05, 0.2);
let mut s = BsmState::new(p, OptionType::Call).unwrap();
let d0 = s.greeks().unwrap().delta;
s.set_spot(110.0).unwrap();
let d1 = s.greeks().unwrap().delta;
assert!(d1 > d0);
}
#[test]
fn set_vol_from_price_round_trip() {
let p = BsmParams::atm_one_year(100.0, 0.05, 0.22);
let mut s = BsmState::new(p, OptionType::Call).unwrap();
let px = s.price().unwrap();
s.set_vol(0.10).unwrap();
let iv = s.set_vol_from_price(px).unwrap();
assert!((iv - 0.22).abs() < 1e-5);
}
}