1use crate::error::{Error, Result};
4use crate::indicators::atr::Atr;
5use crate::ohlcv::Candle;
6use crate::traits::Indicator;
7
8#[derive(Debug, Clone, Copy, PartialEq)]
10pub struct SuperTrendOutput {
11 pub value: f64,
13 pub direction: f64,
16}
17
18#[derive(Debug, Clone, Copy)]
20struct PrevState {
21 final_upper: f64,
22 final_lower: f64,
23 close: f64,
24 direction: f64,
25}
26
27#[derive(Debug, Clone)]
67pub struct SuperTrend {
68 atr: Atr,
69 multiplier: f64,
70 atr_period: usize,
71 prev: Option<PrevState>,
72}
73
74impl SuperTrend {
75 pub fn new(atr_period: usize, multiplier: f64) -> Result<Self> {
82 if !multiplier.is_finite() || multiplier <= 0.0 {
83 return Err(Error::NonPositiveMultiplier);
84 }
85 Ok(Self {
86 atr: Atr::new(atr_period)?,
87 multiplier,
88 atr_period,
89 prev: None,
90 })
91 }
92
93 pub fn classic() -> Self {
95 Self::new(10, 3.0).expect("classic SuperTrend params are valid")
96 }
97
98 pub const fn params(&self) -> (usize, f64) {
100 (self.atr_period, self.multiplier)
101 }
102}
103
104impl Indicator for SuperTrend {
105 type Input = Candle;
106 type Output = SuperTrendOutput;
107
108 fn update(&mut self, candle: Candle) -> Option<SuperTrendOutput> {
109 let atr = self.atr.update(candle)?;
110 let hl2 = f64::midpoint(candle.high, candle.low);
111 let basic_upper = hl2 + self.multiplier * atr;
112 let basic_lower = hl2 - self.multiplier * atr;
113
114 let (final_upper, final_lower, direction) = match self.prev {
115 None => {
116 (basic_upper, basic_lower, 1.0)
118 }
119 Some(p) => {
120 let final_upper = if basic_upper < p.final_upper || p.close > p.final_upper {
121 basic_upper
122 } else {
123 p.final_upper
124 };
125 let final_lower = if basic_lower > p.final_lower || p.close < p.final_lower {
126 basic_lower
127 } else {
128 p.final_lower
129 };
130 let direction = if p.direction < 0.0 {
131 if candle.close <= final_upper {
133 -1.0
134 } else {
135 1.0
136 }
137 } else {
138 if candle.close >= final_lower {
140 1.0
141 } else {
142 -1.0
143 }
144 };
145 (final_upper, final_lower, direction)
146 }
147 };
148
149 let value = if direction > 0.0 {
150 final_lower
151 } else {
152 final_upper
153 };
154 self.prev = Some(PrevState {
155 final_upper,
156 final_lower,
157 close: candle.close,
158 direction,
159 });
160 Some(SuperTrendOutput { value, direction })
161 }
162
163 fn reset(&mut self) {
164 self.atr.reset();
165 self.prev = None;
166 }
167
168 #[inline]
169 fn warmup_period(&self) -> usize {
170 self.atr_period
171 }
172
173 #[inline]
174 fn is_ready(&self) -> bool {
175 self.prev.is_some()
176 }
177
178 #[inline]
179 fn name(&self) -> &'static str {
180 "SuperTrend"
181 }
182}
183
184#[cfg(test)]
185mod tests {
186 use super::*;
187 use crate::traits::BatchExt;
188
189 fn c(high: f64, low: f64, close: f64, ts: i64) -> Candle {
190 Candle::new(f64::midpoint(high, low), high, low, close, 1.0, ts).unwrap()
191 }
192
193 #[test]
194 fn uptrend_keeps_line_below_price_and_direction_up() {
195 let candles: Vec<Candle> = (0..60)
196 .map(|i| {
197 let base = 100.0 + 2.0 * i as f64;
198 c(base + 1.0, base - 1.0, base + 0.5, i)
199 })
200 .collect();
201 let mut st = SuperTrend::classic();
202 for (o, candle) in st.batch(&candles).into_iter().zip(candles.iter()) {
203 if let Some(o) = o {
204 assert_eq!(o.direction, 1.0, "a pure uptrend stays in direction +1");
205 assert!(o.value < candle.close, "the stop line sits below price");
206 }
207 }
208 }
209
210 #[test]
211 fn downtrend_keeps_line_above_price_and_direction_down() {
212 let candles: Vec<Candle> = (0..60)
213 .map(|i| {
214 let base = 220.0 - 2.0 * i as f64;
215 c(base + 1.0, base - 1.0, base - 0.5, i)
216 })
217 .collect();
218 let mut st = SuperTrend::classic();
219 let emitted: Vec<(SuperTrendOutput, f64)> = st
220 .batch(&candles)
221 .into_iter()
222 .zip(candles.iter())
223 .filter_map(|(o, c)| o.map(|v| (v, c.close)))
224 .collect();
225 for &(o, close) in emitted.iter().skip(10) {
228 assert_eq!(
229 o.direction, -1.0,
230 "a steep downtrend settles to direction -1"
231 );
232 assert!(o.value > close, "the stop line sits above price");
233 }
234 }
235
236 #[test]
237 fn trend_flips_when_price_reverses() {
238 let mut candles: Vec<Candle> = (0..40)
239 .map(|i| {
240 let base = 100.0 + i as f64;
241 c(base + 1.0, base - 1.0, base + 0.5, i)
242 })
243 .collect();
244 candles.extend((0..40).map(|i| {
245 let base = 140.0 - i as f64;
246 c(base + 1.0, base - 1.0, base - 0.5, 40 + i)
247 }));
248 let mut st = SuperTrend::classic();
249 let dirs: Vec<f64> = st
250 .batch(&candles)
251 .into_iter()
252 .flatten()
253 .map(|o| o.direction)
254 .collect();
255 assert!(dirs.iter().any(|&d| d > 0.0), "expected an uptrend stretch");
256 assert!(
257 dirs.iter().any(|&d| d < 0.0),
258 "expected a downtrend stretch"
259 );
260 }
261
262 #[test]
263 fn first_emission_matches_warmup_period() {
264 let candles: Vec<Candle> = (0..30)
265 .map(|i| {
266 let base = 100.0 + i as f64;
267 c(base + 1.0, base - 1.0, base, i)
268 })
269 .collect();
270 let mut st = SuperTrend::classic();
271 let out = st.batch(&candles);
272 assert_eq!(st.warmup_period(), 10);
273 for (i, v) in out.iter().enumerate().take(9) {
274 assert!(v.is_none(), "index {i} must be None during warmup");
275 }
276 assert!(out[9].is_some(), "first value lands at warmup_period - 1");
277 }
278
279 #[test]
280 fn rejects_invalid_params() {
281 assert!(SuperTrend::new(0, 3.0).is_err());
282 assert!(SuperTrend::new(10, 0.0).is_err());
283 assert!(SuperTrend::new(10, -1.0).is_err());
284 assert!(SuperTrend::new(10, f64::NAN).is_err());
285 }
286
287 #[test]
290 fn accessors_and_metadata() {
291 let st = SuperTrend::new(10, 3.0).unwrap();
292 let (p, m) = st.params();
293 assert_eq!(p, 10);
294 assert!((m - 3.0).abs() < 1e-12);
295 assert_eq!(st.name(), "SuperTrend");
296 }
297
298 #[test]
299 fn reset_clears_state() {
300 let candles: Vec<Candle> = (0..40)
301 .map(|i| {
302 let base = 100.0 + i as f64;
303 c(base + 1.0, base - 1.0, base, i)
304 })
305 .collect();
306 let mut st = SuperTrend::classic();
307 st.batch(&candles);
308 assert!(st.is_ready());
309 st.reset();
310 assert!(!st.is_ready());
311 assert_eq!(st.update(candles[0]), None);
312 }
313
314 #[test]
315 fn batch_equals_streaming() {
316 let candles: Vec<Candle> = (0..80)
317 .map(|i| {
318 let mid = 100.0 + (i as f64 * 0.3).sin() * 8.0;
319 c(mid + 1.5, mid - 1.5, mid + 0.5, i)
320 })
321 .collect();
322 let mut a = SuperTrend::classic();
323 let mut b = SuperTrend::classic();
324 assert_eq!(
325 a.batch(&candles),
326 candles.iter().map(|x| b.update(*x)).collect::<Vec<_>>()
327 );
328 }
329}