use std::num::NonZeroU32;
use crate::YieldCurveError;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Compounding {
Continuous,
Periodic(NonZeroU32),
Simple,
}
impl Compounding {
#[must_use]
pub fn annual() -> Self {
Self::Periodic(NonZeroU32::new(1).expect("1 is non-zero"))
}
#[must_use]
pub fn semi_annual() -> Self {
Self::Periodic(NonZeroU32::new(2).expect("2 is non-zero"))
}
}
#[must_use]
pub fn discount_factor(rate: f64, t_years: f64, comp: Compounding) -> f64 {
if !rate.is_finite() || !t_years.is_finite() {
return f64::NAN;
}
if t_years == 0.0 {
return 1.0;
}
match comp {
Compounding::Continuous => (-rate * t_years).exp(),
Compounding::Periodic(n) => {
let n = f64::from(n.get());
(1.0 + rate / n).powf(-n * t_years)
}
Compounding::Simple => 1.0 / (1.0 + rate * t_years),
}
}
pub fn forward_rate(
r1: f64,
t1: f64,
r2: f64,
t2: f64,
comp: Compounding,
) -> Result<f64, YieldCurveError> {
for (label, v) in [("r1", r1), ("t1", t1), ("r2", r2), ("t2", t2)] {
if !v.is_finite() {
return Err(YieldCurveError::InvalidTimeRange(format!(
"{label} is not finite ({v})"
)));
}
}
if t1 < 0.0 || t2 < 0.0 {
return Err(YieldCurveError::InvalidTimeRange(format!(
"negative time (t1={t1}, t2={t2})"
)));
}
if t1 >= t2 {
return Err(YieldCurveError::InvalidTimeRange(format!(
"t1 must be < t2 (t1={t1}, t2={t2})"
)));
}
let dt = t2 - t1;
let result = match comp {
Compounding::Continuous => (r2 * t2 - r1 * t1) / dt,
Compounding::Periodic(n) => {
let n = f64::from(n.get());
let num = (1.0 + r2 / n).powf(n * t2);
let den = (1.0 + r1 / n).powf(n * t1);
let ratio = num / den;
n * (ratio.powf(1.0 / (n * dt)) - 1.0)
}
Compounding::Simple => {
let df1 = 1.0 / (1.0 + r1 * t1);
let df2 = 1.0 / (1.0 + r2 * t2);
(df1 / df2 - 1.0) / dt
}
};
if !result.is_finite() {
return Err(YieldCurveError::InvalidTimeRange(format!(
"forward rate is non-finite (r1={r1}, t1={t1}, r2={r2}, t2={t2}, comp={comp:?})"
)));
}
Ok(result)
}
#[cfg(test)]
mod tests {
use super::*;
fn approx_eq(a: f64, b: f64, eps: f64) -> bool {
(a - b).abs() < eps
}
#[test]
fn df_continuous_zero_rate_is_one() {
assert!(approx_eq(
discount_factor(0.0, 5.0, Compounding::Continuous),
1.0,
1e-12
));
}
#[test]
fn df_continuous_known_value() {
assert!(approx_eq(
discount_factor(0.05, 1.0, Compounding::Continuous),
(-0.05_f64).exp(),
1e-12
));
}
#[test]
fn df_annual_known_value() {
assert!(approx_eq(
discount_factor(0.05, 2.0, Compounding::annual()),
1.0 / 1.05_f64.powi(2),
1e-12
));
}
#[test]
fn df_semi_annual() {
assert!(approx_eq(
discount_factor(0.06, 1.0, Compounding::semi_annual()),
1.0 / 1.03_f64.powi(2),
1e-12
));
}
#[test]
fn df_simple() {
assert!(approx_eq(
discount_factor(0.10, 0.5, Compounding::Simple),
1.0 / 1.05,
1e-12
));
}
#[test]
fn df_t_zero_is_one() {
assert_eq!(discount_factor(0.5, 0.0, Compounding::Continuous), 1.0);
assert_eq!(discount_factor(0.5, 0.0, Compounding::annual()), 1.0);
assert_eq!(discount_factor(0.5, 0.0, Compounding::Simple), 1.0);
}
#[test]
fn df_propagates_nan() {
assert!(discount_factor(f64::NAN, 1.0, Compounding::Continuous).is_nan());
assert!(discount_factor(0.05, f64::INFINITY, Compounding::annual()).is_nan());
}
#[test]
fn forward_continuous_classic() {
let f = forward_rate(0.05, 1.0, 0.06, 2.0, Compounding::Continuous).unwrap();
assert!(approx_eq(f, 0.07, 1e-12));
}
#[test]
fn forward_annual_inverse_of_df() {
let f = forward_rate(0.05, 1.0, 0.06, 2.0, Compounding::annual()).unwrap();
let expected = 1.06_f64.powi(2) / 1.05 - 1.0;
assert!(approx_eq(f, expected, 1e-12));
}
#[test]
fn forward_simple_smoke() {
let f = forward_rate(0.05, 0.5, 0.06, 1.0, Compounding::Simple).unwrap();
assert!(f.is_finite());
assert!(f > 0.0);
}
#[test]
fn forward_rejects_t1_ge_t2() {
let err = forward_rate(0.05, 2.0, 0.06, 1.0, Compounding::Continuous).unwrap_err();
assert!(matches!(err, YieldCurveError::InvalidTimeRange(_)));
let err = forward_rate(0.05, 1.0, 0.06, 1.0, Compounding::Continuous).unwrap_err();
assert!(matches!(err, YieldCurveError::InvalidTimeRange(_)));
}
#[test]
fn forward_rejects_negative_time() {
let err = forward_rate(0.05, -0.5, 0.06, 1.0, Compounding::Continuous).unwrap_err();
assert!(matches!(err, YieldCurveError::InvalidTimeRange(_)));
}
#[test]
fn forward_rejects_nan() {
let err = forward_rate(f64::NAN, 1.0, 0.06, 2.0, Compounding::Continuous).unwrap_err();
assert!(matches!(err, YieldCurveError::InvalidTimeRange(_)));
}
}