use crate::stocks::ta::common::opt_cell;
use crate::util::error::{require_finite, FinanceError, FinanceResult};
use crate::util::primitives::PeriodLength;
use crate::{columns_with_strings, print_table_locale_opt};
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub struct StochasticParams {
pub k_period: usize,
pub k_smooth: usize,
pub d_period: usize,
}
impl StochasticParams {
pub const fn fast(k_period: usize, d_period: usize) -> Self {
Self {
k_period,
k_smooth: 1,
d_period,
}
}
pub const fn full(k_period: usize, k_smooth: usize, d_period: usize) -> Self {
Self {
k_period,
k_smooth,
d_period,
}
}
pub const fn warm_up_bars(self) -> usize {
self.k_period
.saturating_add(self.k_smooth.saturating_sub(1))
.saturating_add(self.d_period.saturating_sub(1))
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub struct ValidatedStochastic {
params: StochasticParams,
}
impl ValidatedStochastic {
pub fn new(params: StochasticParams) -> FinanceResult<Self> {
PeriodLength::new(params.k_period)?;
PeriodLength::new(params.k_smooth)?;
PeriodLength::new(params.d_period)?;
Ok(Self { params })
}
#[inline]
pub fn params(self) -> StochasticParams {
self.params
}
pub fn compute(
self,
high: &[f64],
low: &[f64],
close: &[f64],
) -> FinanceResult<StochasticSeries> {
stochastics_validated(high, low, close, self)
}
}
#[derive(Clone, Debug, PartialEq)]
pub struct StochasticSeries {
pub k: Vec<Option<f64>>,
pub d: Vec<Option<f64>>,
pub params: StochasticParams,
}
impl StochasticSeries {
pub fn last_kd(&self) -> Option<(f64, f64)> {
let k = self.k.iter().rev().find_map(|x| *x)?;
let d = self.d.iter().rev().find_map(|x| *x)?;
Some((k, d))
}
}
pub fn stochastics(
high: &[f64],
low: &[f64],
close: &[f64],
params: StochasticParams,
) -> FinanceResult<StochasticSeries> {
let v = ValidatedStochastic::new(params)?;
stochastics_validated(high, low, close, v)
}
pub fn stochastics_solution(
high: &[f64],
low: &[f64],
close: &[f64],
params: StochasticParams,
) -> FinanceResult<StochasticSolution> {
let series = stochastics(high, low, close, params)?;
let formula = format!(
"%K: stoch(k={}, smooth={}); %D: SMA(%K, {})",
params.k_period, params.k_smooth, params.d_period
);
let symbolic =
"raw_%K = 100 * (C - LL) / (HH - LL); %K = SMA(raw_%K, k_smooth); %D = SMA(%K, d)"
.to_string();
Ok(StochasticSolution {
series,
close: close.to_vec(),
formula,
symbolic_formula: symbolic,
})
}
#[derive(Clone, Debug)]
pub struct StochasticSolution {
series: StochasticSeries,
close: Vec<f64>,
formula: String,
symbolic_formula: String,
}
impl StochasticSolution {
pub fn series(&self) -> &StochasticSeries {
&self.series
}
pub fn formula(&self) -> &str {
&self.formula
}
pub fn symbolic_formula(&self) -> &str {
&self.symbolic_formula
}
pub fn params(&self) -> StochasticParams {
self.series.params
}
pub fn print_table(&self) {
self.print_table_locale_opt(None, None);
}
pub fn print_table_locale(&self, locale: &num_format::Locale, precision: usize) {
self.print_table_locale_opt(Some(locale), Some(precision));
}
fn print_table_locale_opt(
&self,
locale: Option<&num_format::Locale>,
precision: Option<usize>,
) {
let columns = columns_with_strings(&[
("period", "i", true),
("close", "f", true),
("k", "f", true),
("d", "f", true),
]);
let data = self
.close
.iter()
.enumerate()
.map(|(i, c)| {
vec![
i.to_string(),
c.to_string(),
opt_cell(self.series.k[i]),
opt_cell(self.series.d[i]),
]
})
.collect();
print_table_locale_opt(&columns, data, locale, precision);
}
}
fn stochastics_validated(
high: &[f64],
low: &[f64],
close: &[f64],
v: ValidatedStochastic,
) -> FinanceResult<StochasticSeries> {
let p = v.params;
check_hlc(high, low, close)?;
let n = close.len();
let mut raw_k = vec![None; n];
let kp = p.k_period;
let mut prev_raw: Option<f64> = None;
for i in 0..n {
if i + 1 < kp {
continue;
}
let start = i + 1 - kp;
let mut hh = f64::NEG_INFINITY;
let mut ll = f64::INFINITY;
for j in start..=i {
hh = hh.max(high[j]);
ll = ll.min(low[j]);
}
let range = hh - ll;
let raw = if range == 0.0 {
prev_raw.unwrap_or(50.0)
} else {
100.0 * (close[i] - ll) / range
};
prev_raw = Some(raw);
raw_k[i] = Some(raw);
}
let smooth_k = sma_option_series(&raw_k, p.k_smooth);
let d_line = sma_option_series(&smooth_k, p.d_period);
Ok(StochasticSeries {
k: smooth_k,
d: d_line,
params: p,
})
}
fn sma_option_series(data: &[Option<f64>], period: usize) -> Vec<Option<f64>> {
let n = data.len();
let mut out = vec![None; n];
if period == 0 || n < period {
return out;
}
for i in (period - 1)..n {
let start = i + 1 - period;
let mut sum = 0.0;
let mut ok = true;
for j in start..=i {
match data[j] {
Some(v) => sum += v,
None => {
ok = false;
break;
}
}
}
if ok {
out[i] = Some(sum / period as f64);
}
}
out
}
fn check_hlc(high: &[f64], low: &[f64], close: &[f64]) -> FinanceResult<()> {
if high.is_empty() {
return Err(FinanceError::EmptyInput { what: "high" });
}
if high.len() != low.len() || high.len() != close.len() {
return Err(FinanceError::LengthMismatch {
left: high.len(),
right: close.len(),
context: "stochastic high/low/close",
});
}
for i in 0..high.len() {
require_finite("high", high[i])?;
require_finite("low", low[i])?;
require_finite("close", close[i])?;
if high[i] < low[i] {
return Err(FinanceError::InvalidCashflow {
message: "high must be >= low for each bar",
});
}
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn fast_const_and_validate() {
let p = StochasticParams::fast(9, 3);
assert_eq!(p.k_smooth, 1);
let v = ValidatedStochastic::new(p).unwrap();
assert_eq!(v.params().k_period, 9);
}
#[test]
fn full_presets() {
let p = StochasticParams::full(14, 3, 3);
assert_eq!(p.warm_up_bars(), 14 + 2 + 2);
}
#[test]
fn series_length_and_warmup() {
let n = 30;
let high: Vec<_> = (0..n).map(|i| 100.0 + i as f64).collect();
let low: Vec<_> = (0..n).map(|i| 90.0 + i as f64).collect();
let close: Vec<_> = (0..n).map(|i| 95.0 + i as f64).collect();
let out = stochastics(&high, &low, &close, StochasticParams::fast(14, 3)).unwrap();
assert_eq!(out.k.len(), n);
assert!(out.k[12].is_none()); assert!(out.k[13].is_some());
assert!(out.d[13 + 2].is_some());
}
#[test]
fn zero_period_err() {
assert!(ValidatedStochastic::new(StochasticParams {
k_period: 0,
k_smooth: 1,
d_period: 3
})
.is_err());
}
#[test]
fn flat_window_carries_previous_raw() {
let high = [10.0, 11.0, 12.0, 12.0, 12.0];
let low = [9.0, 10.0, 12.0, 12.0, 12.0];
let close = [9.5, 10.5, 12.0, 12.0, 12.0];
let p = StochasticParams::fast(3, 1);
let s = stochastics(&high, &low, &close, p).unwrap();
let k3 = s.k[3].unwrap();
assert!((s.k[4].unwrap() - k3).abs() < 1e-12);
assert!(s.k[4].is_some());
}
#[test]
fn k_in_unit_interval_when_range_positive() {
let n = 40;
let high: Vec<_> = (0..n).map(|i| 100.0 + (i % 5) as f64).collect();
let low: Vec<_> = (0..n).map(|i| 90.0 + (i % 5) as f64).collect();
let close: Vec<_> = (0..n).map(|i| 95.0 + (i % 5) as f64 * 0.5).collect();
let s = stochastics(&high, &low, &close, StochasticParams::full(14, 3, 3)).unwrap();
for k in s.k.iter().flatten() {
assert!(*k >= -1e-9 && *k <= 100.0 + 1e-9, "k={k}");
}
}
#[test]
fn high_lt_low_err() {
let h = [10.0, 9.0];
let l = [9.0, 10.0];
let c = [9.5, 9.5];
assert!(stochastics(&h, &l, &c, StochasticParams::fast(2, 1)).is_err());
}
}