1use crate::error::{Error, Result};
4use crate::ohlcv::Candle;
5use crate::traits::Indicator;
6
7#[derive(Debug, Clone, Copy, PartialEq)]
9pub struct VwapStdDevBandsOutput {
10 pub upper: f64,
12 pub middle: f64,
14 pub lower: f64,
16 pub stddev: f64,
18}
19
20#[derive(Debug, Clone)]
54pub struct VwapStdDevBands {
55 multiplier: f64,
56 reference: f64,
61 seeded: bool,
63 sum_dv: f64,
64 sum_d2v: f64,
65 sum_v: f64,
66 has_emitted: bool,
67}
68
69impl VwapStdDevBands {
70 pub fn new(multiplier: f64) -> Result<Self> {
74 if !multiplier.is_finite() || multiplier <= 0.0 {
75 return Err(Error::NonPositiveMultiplier);
76 }
77 Ok(Self {
78 multiplier,
79 reference: 0.0,
80 seeded: false,
81 sum_dv: 0.0,
82 sum_d2v: 0.0,
83 sum_v: 0.0,
84 has_emitted: false,
85 })
86 }
87
88 pub const fn multiplier(&self) -> f64 {
90 self.multiplier
91 }
92}
93
94impl Indicator for VwapStdDevBands {
95 type Input = Candle;
96 type Output = VwapStdDevBandsOutput;
97
98 #[inline]
99 fn update(&mut self, candle: Candle) -> Option<VwapStdDevBandsOutput> {
100 let tp = candle.typical_price();
101 if !self.seeded {
102 self.reference = tp;
103 self.seeded = true;
104 }
105 let d = tp - self.reference;
106 self.sum_dv += d * candle.volume;
107 self.sum_d2v += d * d * candle.volume;
108 self.sum_v += candle.volume;
109 if self.sum_v == 0.0 {
110 return None;
111 }
112 self.has_emitted = true;
113 let mean_d = self.sum_dv / self.sum_v;
116 let vwap = self.reference + mean_d;
117 let var = (self.sum_d2v / self.sum_v - mean_d * mean_d).max(0.0);
120 let sigma = var.sqrt();
121 Some(VwapStdDevBandsOutput {
122 upper: vwap + self.multiplier * sigma,
123 middle: vwap,
124 lower: vwap - self.multiplier * sigma,
125 stddev: sigma,
126 })
127 }
128
129 fn reset(&mut self) {
130 self.reference = 0.0;
131 self.seeded = false;
132 self.sum_dv = 0.0;
133 self.sum_d2v = 0.0;
134 self.sum_v = 0.0;
135 self.has_emitted = false;
136 }
137
138 #[inline]
139 fn warmup_period(&self) -> usize {
140 1
141 }
142
143 #[inline]
144 fn is_ready(&self) -> bool {
145 self.has_emitted
146 }
147
148 #[inline]
149 fn name(&self) -> &'static str {
150 "VwapStdDevBands"
151 }
152}
153
154#[cfg(test)]
155mod tests {
156 use super::*;
157 use crate::traits::BatchExt;
158 use approx::assert_relative_eq;
159
160 fn c(h: f64, l: f64, cl: f64, v: f64) -> Candle {
161 Candle::new(cl, h, l, cl, v, 0).unwrap()
162 }
163
164 #[test]
165 fn rejects_non_positive_multiplier() {
166 assert!(matches!(
167 VwapStdDevBands::new(0.0),
168 Err(Error::NonPositiveMultiplier)
169 ));
170 assert!(matches!(
171 VwapStdDevBands::new(-1.0),
172 Err(Error::NonPositiveMultiplier)
173 ));
174 assert!(matches!(
175 VwapStdDevBands::new(f64::NAN),
176 Err(Error::NonPositiveMultiplier)
177 ));
178 }
179
180 #[test]
181 fn accessors_and_metadata() {
182 let v = VwapStdDevBands::new(2.0).unwrap();
183 assert_relative_eq!(v.multiplier(), 2.0, epsilon = 1e-12);
184 assert_eq!(v.warmup_period(), 1);
185 assert_eq!(v.name(), "VwapStdDevBands");
186 }
187
188 #[test]
189 fn zero_volume_returns_none() {
190 let mut v = VwapStdDevBands::new(2.0).unwrap();
191 assert!(v.update(c(10.0, 10.0, 10.0, 0.0)).is_none());
192 }
193
194 #[test]
195 fn constant_price_collapses_bands() {
196 let candles: Vec<Candle> = (0..10).map(|_| c(10.0, 10.0, 10.0, 5.0)).collect();
197 let mut v = VwapStdDevBands::new(2.0).unwrap();
198 let last = v.batch(&candles).into_iter().flatten().last().unwrap();
199 assert_relative_eq!(last.middle, 10.0, epsilon = 1e-9);
200 assert_relative_eq!(last.stddev, 0.0, epsilon = 1e-9);
201 assert_relative_eq!(last.upper, 10.0, epsilon = 1e-9);
202 assert_relative_eq!(last.lower, 10.0, epsilon = 1e-9);
203 }
204
205 #[test]
206 fn upper_above_middle_above_lower() {
207 let candles: Vec<Candle> = (0..50)
208 .map(|i| {
209 let m = 100.0 + (f64::from(i) * 0.2).sin() * 5.0;
210 c(m + 1.0, m - 1.0, m, 1.0 + f64::from(i % 5))
211 })
212 .collect();
213 let mut v = VwapStdDevBands::new(2.0).unwrap();
214 for o in v.batch(&candles).into_iter().flatten() {
215 assert!(o.upper >= o.middle);
216 assert!(o.middle >= o.lower);
217 assert!(o.stddev >= 0.0);
218 }
219 }
220
221 #[test]
222 fn batch_equals_streaming() {
223 let candles: Vec<Candle> = (0..40)
224 .map(|i| {
225 c(
226 f64::from(i) + 2.0,
227 f64::from(i),
228 f64::from(i) + 1.0,
229 1.0 + f64::from(i % 4),
230 )
231 })
232 .collect();
233 let mut a = VwapStdDevBands::new(2.0).unwrap();
234 let mut b = VwapStdDevBands::new(2.0).unwrap();
235 assert_eq!(
236 a.batch(&candles),
237 candles.iter().map(|x| b.update(*x)).collect::<Vec<_>>()
238 );
239 }
240
241 #[test]
242 fn reset_clears_state() {
243 let candles: Vec<Candle> = (0..10)
244 .map(|i| c(f64::from(i) + 1.0, f64::from(i) - 1.0, f64::from(i), 1.0))
245 .collect();
246 let mut v = VwapStdDevBands::new(2.0).unwrap();
247 v.batch(&candles);
248 assert!(v.is_ready());
249 v.reset();
250 assert!(!v.is_ready());
251 assert_eq!(v.update(c(10.0, 10.0, 10.0, 0.0)), None);
254 }
255
256 #[test]
261 fn reference_values() {
262 let candles = [c(8.0, 8.0, 8.0, 1.0), c(12.0, 12.0, 12.0, 1.0)];
266 let mut v = VwapStdDevBands::new(1.5).unwrap();
267 let _ = v.update(candles[0]);
268 let out = v.update(candles[1]).unwrap();
269 assert_relative_eq!(out.middle, 10.0, epsilon = 1e-9);
270 assert_relative_eq!(out.stddev, 2.0, epsilon = 1e-9);
271 assert_relative_eq!(out.upper, 13.0, epsilon = 1e-9);
272 assert_relative_eq!(out.lower, 7.0, epsilon = 1e-9);
273 }
274
275 #[test]
281 fn deviation_at_a_high_price_level_matches_a_two_pass_reference() {
282 let closes: Vec<f64> = (0..400)
283 .map(|i| {
284 let t = f64::from(i);
285 1e8 + ((t * 0.11).sin() + 0.4 * (t * 0.37).cos())
286 })
287 .collect();
288
289 let mut ind = VwapStdDevBands::new(2.0).unwrap();
290 let (mut prices, mut volumes): (Vec<f64>, Vec<f64>) = (Vec::new(), Vec::new());
291 let mut compared = 0_usize;
292 for (i, &c) in closes.iter().enumerate() {
293 let volume = 10.0 + (i % 7) as f64;
294 let timestamp = i64::try_from(i).unwrap();
297 let candle = Candle::new_unchecked(c, c + 0.5, c - 0.5, c, volume, timestamp);
298 let out = ind.update(candle);
299 prices.push(candle.typical_price());
300 volumes.push(volume);
301 let Some(out) = out else { continue };
302
303 let total: f64 = volumes.iter().sum();
306 let vwap: f64 = prices.iter().zip(&volumes).map(|(p, v)| p * v).sum::<f64>() / total;
307 let var: f64 = prices
308 .iter()
309 .zip(&volumes)
310 .map(|(p, v)| v * (p - vwap) * (p - vwap))
311 .sum::<f64>()
312 / total;
313 compared += 1;
314 assert_relative_eq!(out.middle, vwap, max_relative = 1e-14);
315 assert_relative_eq!(out.stddev, var.sqrt(), max_relative = 1e-9);
316 }
317 assert_eq!(compared, closes.len());
318 }
319}