Skip to main content

fin_primitives/signals/
compose.rs

1//! # Signal Composition Engine
2//!
3//! A composable, expression-tree DSL for building derived signals from existing
4//! indicators. Instead of writing bespoke structs for every derived computation,
5//! callers use [`SignalExpr`] to describe *what* to compute and [`ComposedSignal`]
6//! to evaluate that expression on each new bar.
7//!
8//! ## Architecture
9//!
10//! ```text
11//! SignalExpr (description)
12//!     └── ComposedSignal (stateful evaluator)
13//!             └── inner Signal impls (Sma, Ema, Rsi, …)
14//! ```
15//!
16//! [`SignalExpr`] is a pure data structure — it describes the computation graph
17//! without owning any mutable indicator state. [`ComposedSignal`] is the stateful
18//! counterpart that actually holds the inner indicators and evaluates the tree.
19//!
20//! ## Warmup Semantics
21//!
22//! Composed signals correctly propagate warmup: if any leaf signal in the expression
23//! tree is not yet ready, the composed signal returns `SignalValue::Unavailable`.
24//! The warmup period of a composition is the *maximum* warmup period of all leaves,
25//! plus any additional lag introduced by [`SignalExpr::Lag`] nodes.
26//!
27//! ## Builder API
28//!
29//! Use [`SignalBuilder`] for a fluent, method-chaining API. The builder works on
30//! any concrete `Signal` type and produces a [`ComposedSignal`] ready for use.
31//!
32//! ## Example
33//!
34//! ```rust
35//! use fin_primitives::signals::indicators::Rsi;
36//! use fin_primitives::signals::{BarInput, Signal};
37//! use fin_primitives::signals::compose::{SignalBuilder, NormMethod, Direction};
38//! use rust_decimal_macros::dec;
39//!
40//! let rsi = Rsi::new("rsi14", 14).unwrap();
41//!
42//! // Build: RSI(14) → lag(1) → normalize(ZScore) → threshold(2.0, Above)
43//! let mut composed = SignalBuilder::new(rsi)
44//!     .lag(1)
45//!     .normalize(NormMethod::ZScore)
46//!     .threshold(dec!(2), Direction::Above)
47//!     .build();
48//!
49//! let bar = BarInput::from_close(dec!(100));
50//! let _value = composed.update(&bar); // Ok(Unavailable) during warmup
51//! ```
52
53use crate::error::FinError;
54use crate::signals::{BarInput, Signal, SignalValue};
55use rust_decimal::Decimal;
56use std::collections::VecDeque;
57
58// ── NormMethod ────────────────────────────────────────────────────────────────
59
60/// Method used to normalise a signal's output stream.
61#[derive(Debug, Clone, Copy, PartialEq, Eq)]
62pub enum NormMethod {
63    /// Scales the value to `[0, 1]` using the rolling min/max over the last `window` bars.
64    ///
65    /// `output = (value - min) / (max - min)`
66    MinMax,
67
68    /// Standardises to zero mean and unit variance over the last `window` bars.
69    ///
70    /// `output = (value - mean) / std_dev`
71    ZScore,
72
73    /// Expresses the value as its percentile rank within the last `window` bars.
74    ///
75    /// `output ∈ [0, 1]` where `1.0` = largest value in the window.
76    Percentile,
77}
78
79// ── Direction ─────────────────────────────────────────────────────────────────
80
81/// Direction of a threshold test applied to a signal.
82#[derive(Debug, Clone, Copy, PartialEq, Eq)]
83pub enum Direction {
84    /// Passes `+1` when the signal is strictly above the threshold, `0` otherwise.
85    Above,
86    /// Passes `-1` when the signal is strictly below the threshold, `0` otherwise.
87    Below,
88    /// Passes `+1` on an upward cross, `-1` on a downward cross, `0` otherwise.
89    Cross,
90}
91
92// ── SignalKind ────────────────────────────────────────────────────────────────
93
94/// Identifies a named leaf signal within a composed expression.
95///
96/// In a [`ComposedSignal`], every `Raw` leaf in the expression tree corresponds
97/// to an entry in the signal registry keyed by `name`.
98#[derive(Debug, Clone)]
99pub struct SignalKind {
100    /// The name of the signal as returned by [`Signal::name`].
101    pub name: String,
102}
103
104impl SignalKind {
105    /// Creates a new `SignalKind` referencing a signal by name.
106    pub fn new(name: impl Into<String>) -> Self {
107        Self { name: name.into() }
108    }
109}
110
111// ── SignalExpr ────────────────────────────────────────────────────────────────
112
113/// A composable expression tree for building derived signals.
114///
115/// `SignalExpr` is a pure description of a computation — it does not hold any
116/// mutable state. Feed it into [`ComposedSignal`] to evaluate it on bar data.
117///
118/// # Expression Nodes
119///
120/// | Variant | Description |
121/// |---------|-------------|
122/// | `Raw` | Leaf: the raw output of a named indicator |
123/// | `Add` | Element-wise addition of two sub-expressions |
124/// | `Sub` | Element-wise subtraction |
125/// | `Mul` | Multiply a sub-expression by a scalar constant |
126/// | `Lag` | Delay a sub-expression by `n` bars |
127/// | `Normalize` | Normalise a sub-expression using a rolling window |
128/// | `Threshold` | Convert a scalar sub-expression to `+1`, `0`, or `-1` |
129#[derive(Debug, Clone)]
130pub enum SignalExpr {
131    /// A raw indicator output, identified by name.
132    Raw(SignalKind),
133
134    /// The sum of two sub-expressions. Returns `Unavailable` if either is.
135    Add(Box<SignalExpr>, Box<SignalExpr>),
136
137    /// The difference of two sub-expressions (`left - right`). Returns `Unavailable` if either is.
138    Sub(Box<SignalExpr>, Box<SignalExpr>),
139
140    /// A sub-expression scaled by a constant factor.
141    Mul(Box<SignalExpr>, Decimal),
142
143    /// A sub-expression delayed by `n` bars.
144    ///
145    /// Returns `Unavailable` until the buffer has accumulated `n` bars of ready values.
146    Lag(Box<SignalExpr>, usize),
147
148    /// A sub-expression normalised using a rolling window of `window` bars.
149    ///
150    /// Returns `Unavailable` until `window` ready values have been accumulated.
151    Normalize(Box<SignalExpr>, NormMethod, usize),
152
153    /// Converts a sub-expression to a directional signal relative to a threshold.
154    ///
155    /// Output is `Scalar(1)`, `Scalar(-1)`, or `Scalar(0)`.
156    Threshold(Box<SignalExpr>, Decimal, Direction),
157}
158
159impl SignalExpr {
160    /// Constructs a `Raw` leaf from a signal name.
161    pub fn raw(name: impl Into<String>) -> Self {
162        Self::Raw(SignalKind::new(name))
163    }
164
165    /// Wraps `self` in an `Add` with `rhs`.
166    pub fn add(self, rhs: SignalExpr) -> Self {
167        Self::Add(Box::new(self), Box::new(rhs))
168    }
169
170    /// Wraps `self` in a `Sub` with `rhs`.
171    pub fn sub(self, rhs: SignalExpr) -> Self {
172        Self::Sub(Box::new(self), Box::new(rhs))
173    }
174
175    /// Wraps `self` in a `Mul` with scalar `factor`.
176    pub fn mul(self, factor: Decimal) -> Self {
177        Self::Mul(Box::new(self), factor)
178    }
179
180    /// Wraps `self` in a `Lag` of `n` bars.
181    pub fn lag(self, n: usize) -> Self {
182        Self::Lag(Box::new(self), n)
183    }
184
185    /// Wraps `self` in a `Normalize` node.
186    pub fn normalize(self, method: NormMethod, window: usize) -> Self {
187        Self::Normalize(Box::new(self), method, window)
188    }
189
190    /// Wraps `self` in a `Threshold` node.
191    pub fn threshold(self, level: Decimal, direction: Direction) -> Self {
192        Self::Threshold(Box::new(self), level, direction)
193    }
194
195    /// Returns a flat list of all leaf signal names referenced by this expression.
196    pub fn leaf_names(&self) -> Vec<&str> {
197        let mut names = Vec::new();
198        self.collect_leaf_names(&mut names);
199        names
200    }
201
202    fn collect_leaf_names<'a>(&'a self, out: &mut Vec<&'a str>) {
203        match self {
204            Self::Raw(kind) => out.push(&kind.name),
205            Self::Add(l, r) | Self::Sub(l, r) => {
206                l.collect_leaf_names(out);
207                r.collect_leaf_names(out);
208            }
209            Self::Mul(inner, _)
210            | Self::Lag(inner, _)
211            | Self::Normalize(inner, _, _)
212            | Self::Threshold(inner, _, _) => inner.collect_leaf_names(out),
213        }
214    }
215}
216
217// ── ExprState ─────────────────────────────────────────────────────────────────
218
219/// Internal mutable state for a single node in the expression tree.
220///
221/// The state tree mirrors the `SignalExpr` tree 1:1 so that each stateful node
222/// (Lag buffer, Normalize rolling window) can be updated independently.
223enum ExprState {
224    Raw,
225    Add(Box<ExprState>, Box<ExprState>),
226    Sub(Box<ExprState>, Box<ExprState>),
227    Mul(Box<ExprState>),
228    Lag {
229        inner: Box<ExprState>,
230        buffer: VecDeque<SignalValue>,
231        n: usize,
232    },
233    Normalize {
234        inner: Box<ExprState>,
235        window: VecDeque<Decimal>,
236        window_size: usize,
237        method: NormMethod,
238        prev: SignalValue,
239    },
240    Threshold {
241        inner: Box<ExprState>,
242        level: Decimal,
243        direction: Direction,
244        prev: SignalValue,
245    },
246}
247
248impl ExprState {
249    /// Builds an `ExprState` tree from a `SignalExpr` tree.
250    fn from_expr(expr: &SignalExpr) -> Self {
251        match expr {
252            SignalExpr::Raw(_) => Self::Raw,
253            SignalExpr::Add(l, r) => {
254                Self::Add(Box::new(Self::from_expr(l)), Box::new(Self::from_expr(r)))
255            }
256            SignalExpr::Sub(l, r) => {
257                Self::Sub(Box::new(Self::from_expr(l)), Box::new(Self::from_expr(r)))
258            }
259            SignalExpr::Mul(inner, _) => Self::Mul(Box::new(Self::from_expr(inner))),
260            SignalExpr::Lag(inner, n) => Self::Lag {
261                inner: Box::new(Self::from_expr(inner)),
262                buffer: VecDeque::new(),
263                n: *n,
264            },
265            SignalExpr::Normalize(inner, method, window_size) => Self::Normalize {
266                inner: Box::new(Self::from_expr(inner)),
267                window: VecDeque::new(),
268                window_size: *window_size,
269                method: *method,
270                prev: SignalValue::Unavailable,
271            },
272            SignalExpr::Threshold(inner, _, direction) => Self::Threshold {
273                inner: Box::new(Self::from_expr(inner)),
274                level: Decimal::ZERO, // overwritten during eval
275                direction: *direction,
276                prev: SignalValue::Unavailable,
277            },
278        }
279    }
280
281    /// Evaluates this state node given the raw signal values for all leaves.
282    fn eval(
283        &mut self,
284        expr: &SignalExpr,
285        leaf_values: &std::collections::HashMap<String, SignalValue>,
286    ) -> SignalValue {
287        match (self, expr) {
288            (Self::Raw, SignalExpr::Raw(kind)) => leaf_values
289                .get(&kind.name)
290                .cloned()
291                .unwrap_or(SignalValue::Unavailable),
292
293            (Self::Add(ls, rs), SignalExpr::Add(le, re)) => {
294                let l = ls.eval(le, leaf_values);
295                let r = rs.eval(re, leaf_values);
296                l.add(r)
297            }
298
299            (Self::Sub(ls, rs), SignalExpr::Sub(le, re)) => {
300                let l = ls.eval(le, leaf_values);
301                let r = rs.eval(re, leaf_values);
302                l.sub(r)
303            }
304
305            (Self::Mul(inner_state), SignalExpr::Mul(inner_expr, factor)) => {
306                let v = inner_state.eval(inner_expr, leaf_values);
307                v.mul(*factor)
308            }
309
310            (
311                Self::Lag { inner, buffer, n },
312                SignalExpr::Lag(inner_expr, _),
313            ) => {
314                let v = inner.eval(inner_expr, leaf_values);
315                if *n == 0 {
316                    return v;
317                }
318                // Push the current value into the lag buffer.
319                buffer.push_back(v);
320                // Return the value that was at position `n` in the past.
321                if buffer.len() > *n {
322                    buffer.pop_front().unwrap_or(SignalValue::Unavailable)
323                } else {
324                    SignalValue::Unavailable
325                }
326            }
327
328            (
329                Self::Normalize { inner, window, window_size, method, .. },
330                SignalExpr::Normalize(inner_expr, _, _),
331            ) => {
332                let v = inner.eval(inner_expr, leaf_values);
333                match v {
334                    SignalValue::Unavailable => SignalValue::Unavailable,
335                    SignalValue::Scalar(d) => {
336                        window.push_back(d);
337                        if window.len() > *window_size {
338                            window.pop_front();
339                        }
340                        if window.len() < *window_size {
341                            return SignalValue::Unavailable;
342                        }
343                        compute_norm(window, *method, d)
344                    }
345                }
346            }
347
348            (
349                Self::Threshold { inner, prev, direction, level },
350                SignalExpr::Threshold(inner_expr, threshold_level, _),
351            ) => {
352                *level = *threshold_level;
353                let v = inner.eval(inner_expr, leaf_values);
354                let result = match direction {
355                    Direction::Above => match &v {
356                        SignalValue::Scalar(curr) if *curr > *level => {
357                            SignalValue::Scalar(Decimal::ONE)
358                        }
359                        SignalValue::Scalar(_) => SignalValue::Scalar(Decimal::ZERO),
360                        SignalValue::Unavailable => SignalValue::Unavailable,
361                    },
362                    Direction::Below => match &v {
363                        SignalValue::Scalar(curr) if *curr < *level => {
364                            SignalValue::Scalar(-Decimal::ONE)
365                        }
366                        SignalValue::Scalar(_) => SignalValue::Scalar(Decimal::ZERO),
367                        SignalValue::Unavailable => SignalValue::Unavailable,
368                    },
369                    Direction::Cross => {
370                        
371                        match (&v, &*prev) {
372                            (SignalValue::Scalar(curr), SignalValue::Scalar(p)) => {
373                                if *curr > *level && *p <= *level {
374                                    SignalValue::Scalar(Decimal::ONE)
375                                } else if *curr < *level && *p >= *level {
376                                    SignalValue::Scalar(-Decimal::ONE)
377                                } else {
378                                    SignalValue::Scalar(Decimal::ZERO)
379                                }
380                            }
381                            _ => SignalValue::Unavailable,
382                        }
383                    }
384                };
385                *prev = v;
386                result
387            }
388
389            // Mismatched arms should never occur if ExprState::from_expr is consistent.
390            _ => SignalValue::Unavailable,
391        }
392    }
393
394    /// Resets all buffered state recursively.
395    fn reset(&mut self) {
396        match self {
397            Self::Raw | Self::Mul(_) => {}
398            Self::Add(l, r) | Self::Sub(l, r) => {
399                l.reset();
400                r.reset();
401            }
402            Self::Lag { inner, buffer, .. } => {
403                inner.reset();
404                buffer.clear();
405            }
406            Self::Normalize { inner, window, prev, .. } => {
407                inner.reset();
408                window.clear();
409                *prev = SignalValue::Unavailable;
410            }
411            Self::Threshold { inner, prev, .. } => {
412                inner.reset();
413                *prev = SignalValue::Unavailable;
414            }
415        }
416    }
417}
418
419/// Computes a normalised value from a rolling window.
420fn compute_norm(
421    window: &VecDeque<Decimal>,
422    method: NormMethod,
423    current: Decimal,
424) -> SignalValue {
425    if window.is_empty() {
426        return SignalValue::Unavailable;
427    }
428    match method {
429        NormMethod::MinMax => {
430            let min = window.iter().copied().fold(current, Decimal::min);
431            let max = window.iter().copied().fold(current, Decimal::max);
432            let range = max - min;
433            if range.is_zero() {
434                SignalValue::Scalar(Decimal::ZERO)
435            } else {
436                match (current - min).checked_div(range) {
437                    Some(v) => SignalValue::Scalar(v),
438                    None => SignalValue::Unavailable,
439                }
440            }
441        }
442        NormMethod::ZScore => {
443            let n = window.len() as f64;
444            if n < 2.0 {
445                return SignalValue::Unavailable;
446            }
447            use rust_decimal::prelude::ToPrimitive;
448            let mean: f64 = window.iter().filter_map(|v| v.to_f64()).sum::<f64>() / n;
449            let variance: f64 = window
450                .iter()
451                .filter_map(|v| v.to_f64())
452                .map(|v| (v - mean).powi(2))
453                .sum::<f64>()
454                / (n - 1.0);
455            let std_dev = variance.sqrt();
456            if std_dev == 0.0 {
457                return SignalValue::Scalar(Decimal::ZERO);
458            }
459            let curr_f = current.to_f64().unwrap_or(mean);
460            match Decimal::try_from((curr_f - mean) / std_dev) {
461                Ok(z) => SignalValue::Scalar(z),
462                Err(_) => SignalValue::Unavailable,
463            }
464        }
465        NormMethod::Percentile => {
466            let n = window.len();
467            let count_below = window.iter().filter(|&&v| v < current).count();
468            let count_equal = window.iter().filter(|&&v| v == current).count();
469            // Percentile rank: (count_below + 0.5 * count_equal) / n
470            let rank_f = (count_below as f64 + 0.5 * count_equal as f64) / n as f64;
471            match Decimal::try_from(rank_f) {
472                Ok(rank) => SignalValue::Scalar(rank),
473                Err(_) => SignalValue::Unavailable,
474            }
475        }
476    }
477}
478
479// ── ComposedSignal ────────────────────────────────────────────────────────────
480
481/// Evaluates a [`SignalExpr`] expression tree on each new bar.
482///
483/// `ComposedSignal` holds a set of leaf [`Signal`] implementations and an
484/// expression tree that describes how to combine them. On each call to
485/// [`Signal::update`], it updates all leaves, then evaluates the expression
486/// tree bottom-up to produce the final output.
487///
488/// The warmup period is the maximum leaf warmup period plus any lag introduced
489/// by [`SignalExpr::Lag`] nodes at the outermost level.
490///
491/// # Construction
492///
493/// Use the [`SignalBuilder`] fluent API for the most ergonomic construction, or
494/// build a [`SignalExpr`] tree manually and pass it to [`ComposedSignal::new`].
495pub struct ComposedSignal {
496    name: String,
497    expr: SignalExpr,
498    state: ExprState,
499    leaves: Vec<Box<dyn Signal>>,
500    bars_seen: usize,
501}
502
503impl ComposedSignal {
504    /// Constructs a `ComposedSignal` from a name, expression tree, and leaf signals.
505    ///
506    /// The `leaves` vector must contain one signal per unique name referenced in the
507    /// expression tree. Signals are looked up by name during evaluation.
508    ///
509    /// # Errors
510    ///
511    /// Returns [`FinError::InvalidInput`] if `leaves` is empty.
512    pub fn new(
513        name: impl Into<String>,
514        expr: SignalExpr,
515        leaves: Vec<Box<dyn Signal>>,
516    ) -> Result<Self, FinError> {
517        if leaves.is_empty() {
518            return Err(FinError::InvalidInput(
519                "ComposedSignal requires at least one leaf signal".into(),
520            ));
521        }
522        let state = ExprState::from_expr(&expr);
523        Ok(Self {
524            name: name.into(),
525            expr,
526            state,
527            leaves,
528            bars_seen: 0,
529        })
530    }
531
532    /// Returns the maximum warmup period across all leaf signals.
533    pub fn leaf_warmup_period(&self) -> usize {
534        self.leaves.iter().map(|s| s.period()).max().unwrap_or(0)
535    }
536}
537
538impl Signal for ComposedSignal {
539    fn name(&self) -> &str {
540        &self.name
541    }
542
543    fn update(&mut self, bar: &BarInput) -> Result<SignalValue, FinError> {
544        self.bars_seen += 1;
545
546        // Update all leaves and collect their values into a name-keyed map.
547        let mut leaf_values = std::collections::HashMap::with_capacity(self.leaves.len());
548        for leaf in &mut self.leaves {
549            let val = leaf.update(bar)?;
550            leaf_values.insert(leaf.name().to_owned(), val);
551        }
552
553        // Evaluate the expression tree using the collected leaf values.
554        let result = self.state.eval(&self.expr, &leaf_values);
555        Ok(result)
556    }
557
558    fn is_ready(&self) -> bool {
559        self.leaves.iter().all(|s| s.is_ready())
560    }
561
562    fn period(&self) -> usize {
563        self.leaf_warmup_period()
564    }
565
566    fn reset(&mut self) {
567        for leaf in &mut self.leaves {
568            leaf.reset();
569        }
570        self.state.reset();
571        self.bars_seen = 0;
572    }
573}
574
575// ── SignalBuilder ─────────────────────────────────────────────────────────────
576
577/// Fluent builder for [`ComposedSignal`] using method-chaining.
578///
579/// Start with any concrete signal type that implements [`Signal`] and chain
580/// transformations. Each method wraps the accumulated expression in a new
581/// [`SignalExpr`] node.
582///
583/// # Example
584///
585/// ```rust
586/// use fin_primitives::signals::indicators::Sma;
587/// use fin_primitives::signals::{BarInput, Signal};
588/// use fin_primitives::signals::compose::{SignalBuilder, NormMethod, Direction};
589/// use rust_decimal_macros::dec;
590///
591/// let sma = Sma::new("sma20", 20).unwrap();
592/// let mut composed = SignalBuilder::new(sma)
593///     .lag(2)
594///     .normalize(NormMethod::MinMax)
595///     .build();
596///
597/// let bar = BarInput::from_close(dec!(100));
598/// let _ = composed.update(&bar);
599/// ```
600pub struct SignalBuilder<S: Signal + 'static> {
601    signal: S,
602    /// The accumulated expression tree (grows as builder methods are called).
603    expr: SignalExpr,
604    /// Window size used for `Normalize` nodes (default: 20).
605    norm_window: usize,
606}
607
608impl<S: Signal + 'static> SignalBuilder<S> {
609    /// Creates a builder from a concrete signal.
610    ///
611    /// The initial expression is `Raw(signal.name())`.
612    pub fn new(signal: S) -> Self {
613        let name = signal.name().to_owned();
614        Self {
615            signal,
616            expr: SignalExpr::raw(name),
617            norm_window: 20,
618        }
619    }
620
621    /// Sets the rolling window size used by subsequent `normalize()` calls (default: 20).
622    pub fn with_norm_window(mut self, window: usize) -> Self {
623        self.norm_window = window;
624        self
625    }
626
627    /// Wraps the current expression in a `Lag` of `n` bars.
628    pub fn lag(mut self, n: usize) -> Self {
629        self.expr = self.expr.lag(n);
630        self
631    }
632
633    /// Wraps the current expression in a `Normalize` node using the configured window.
634    pub fn normalize(mut self, method: NormMethod) -> Self {
635        let window = self.norm_window;
636        self.expr = self.expr.normalize(method, window);
637        self
638    }
639
640    /// Wraps the current expression in a `Normalize` node with an explicit `window`.
641    pub fn normalize_window(mut self, method: NormMethod, window: usize) -> Self {
642        self.expr = self.expr.normalize(method, window);
643        self
644    }
645
646    /// Wraps the current expression in a `Threshold` node.
647    pub fn threshold(mut self, level: Decimal, direction: Direction) -> Self {
648        self.expr = self.expr.threshold(level, direction);
649        self
650    }
651
652    /// Scales the current expression by `factor`.
653    pub fn scale(mut self, factor: Decimal) -> Self {
654        self.expr = self.expr.mul(factor);
655        self
656    }
657
658    /// Consumes the builder and produces a [`ComposedSignal`].
659    ///
660    /// The composed signal name is derived from the inner signal's name.
661    pub fn build(self) -> ComposedSignal {
662        let composed_name = format!("composed({})", self.signal.name());
663        self.build_named(composed_name)
664    }
665
666    /// Consumes the builder and produces a [`ComposedSignal`] with an explicit name.
667    pub fn build_named(self, name: impl Into<String>) -> ComposedSignal {
668        // One leaf, so the emptiness check in `ComposedSignal::new` cannot fail; build
669        // the value directly instead of unwrapping that `Result`.
670        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(self.signal)];
671        let state = ExprState::from_expr(&self.expr);
672        ComposedSignal { name: name.into(), expr: self.expr, state, leaves, bars_seen: 0 }
673    }
674}
675
676// ── Tests ─────────────────────────────────────────────────────────────────────
677
678#[cfg(test)]
679mod tests {
680    use super::*;
681    use crate::signals::indicators::{Ema, Rsi, Sma};
682    use rust_decimal_macros::dec;
683
684    fn bar(close: &str) -> BarInput {
685        BarInput::from_close(close.parse().unwrap())
686    }
687
688    fn feed_n(signal: &mut impl Signal, close: &str, n: usize) {
689        for _ in 0..n {
690            signal.update(&bar(close)).unwrap();
691        }
692    }
693
694    // ── SignalExpr leaf_names ────────────────────────────────────────────────
695
696    #[test]
697    fn test_expr_raw_leaf_name() {
698        let expr = SignalExpr::raw("sma5");
699        assert_eq!(expr.leaf_names(), vec!["sma5"]);
700    }
701
702    #[test]
703    fn test_expr_add_leaf_names() {
704        let expr = SignalExpr::raw("a").add(SignalExpr::raw("b"));
705        let names = expr.leaf_names();
706        assert!(names.contains(&"a"));
707        assert!(names.contains(&"b"));
708    }
709
710    #[test]
711    fn test_expr_nested_leaf_names() {
712        let expr = SignalExpr::raw("sma5")
713            .lag(1)
714            .normalize(NormMethod::ZScore, 20)
715            .threshold(dec!(0), Direction::Above);
716        assert_eq!(expr.leaf_names(), vec!["sma5"]);
717    }
718
719    // ── ComposedSignal: Raw passthrough ──────────────────────────────────────
720
721    #[test]
722    fn test_composed_raw_passthrough() {
723        let sma = Sma::new("sma3", 3).unwrap();
724        let expr = SignalExpr::raw("sma3");
725        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma)];
726        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
727
728        feed_n(&mut composed, "10", 2);
729        let v = composed.update(&bar("10")).unwrap();
730        assert!(matches!(v, SignalValue::Scalar(_)));
731    }
732
733    #[test]
734    fn test_composed_raw_unavailable_during_warmup() {
735        let sma = Sma::new("sma5", 5).unwrap();
736        let expr = SignalExpr::raw("sma5");
737        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma)];
738        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
739
740        let v = composed.update(&bar("10")).unwrap();
741        assert_eq!(v, SignalValue::Unavailable);
742    }
743
744    // ── Mul ──────────────────────────────────────────────────────────────────
745
746    #[test]
747    fn test_composed_mul_scales_value() {
748        let sma = Sma::new("sma1", 1).unwrap();
749        let expr = SignalExpr::raw("sma1").mul(dec!(2));
750        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma)];
751        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
752
753        let v = composed.update(&bar("10")).unwrap();
754        assert_eq!(v, SignalValue::Scalar(dec!(20)));
755    }
756
757    // ── Add ──────────────────────────────────────────────────────────────────
758
759    #[test]
760    fn test_composed_add_two_signals() {
761        let sma1 = Sma::new("sma_a", 1).unwrap();
762        let sma2 = Sma::new("sma_b", 1).unwrap();
763        let expr = SignalExpr::raw("sma_a").add(SignalExpr::raw("sma_b"));
764        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma1), Box::new(sma2)];
765        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
766
767        let v = composed.update(&bar("15")).unwrap();
768        // sma_a = 15, sma_b = 15, sum = 30
769        assert_eq!(v, SignalValue::Scalar(dec!(30)));
770    }
771
772    // ── Lag ──────────────────────────────────────────────────────────────────
773
774    #[test]
775    fn test_composed_lag_delays_values() {
776        let sma = Sma::new("sma1", 1).unwrap();
777        let expr = SignalExpr::raw("sma1").lag(2);
778        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma)];
779        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
780
781        // bars 1,2: buffer fills but has not produced 2 prior values yet
782        composed.update(&bar("10")).unwrap();
783        composed.update(&bar("20")).unwrap();
784        // bar 3: lag-2 should produce bar-1's value = 10
785        let v = composed.update(&bar("30")).unwrap();
786        assert_eq!(v, SignalValue::Scalar(dec!(10)));
787    }
788
789    #[test]
790    fn test_composed_lag_zero_is_passthrough() {
791        let sma = Sma::new("sma1", 1).unwrap();
792        let expr = SignalExpr::raw("sma1").lag(0);
793        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma)];
794        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
795
796        let v = composed.update(&bar("42")).unwrap();
797        assert_eq!(v, SignalValue::Scalar(dec!(42)));
798    }
799
800    // ── Normalize: MinMax ────────────────────────────────────────────────────
801
802    #[test]
803    fn test_normalize_minmax_range_of_constant_returns_zero() {
804        let sma = Sma::new("sma1", 1).unwrap();
805        let expr = SignalExpr::raw("sma1").normalize(NormMethod::MinMax, 3);
806        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma)];
807        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
808
809        // Feed 3 identical values → min==max → output 0
810        composed.update(&bar("10")).unwrap();
811        composed.update(&bar("10")).unwrap();
812        let v = composed.update(&bar("10")).unwrap();
813        assert_eq!(v, SignalValue::Scalar(dec!(0)));
814    }
815
816    #[test]
817    fn test_normalize_minmax_high_value_approaches_one() {
818        let sma = Sma::new("sma1", 1).unwrap();
819        let expr = SignalExpr::raw("sma1").normalize(NormMethod::MinMax, 3);
820        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma)];
821        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
822
823        // window: [0, 50, 100] → current=100, min=0, max=100 → 1.0
824        composed.update(&bar("0")).unwrap();
825        composed.update(&bar("50")).unwrap();
826        let v = composed.update(&bar("100")).unwrap();
827        assert_eq!(v, SignalValue::Scalar(dec!(1)));
828    }
829
830    // ── Normalize: ZScore ────────────────────────────────────────────────────
831
832    #[test]
833    fn test_normalize_zscore_mean_value_near_zero() {
834        let sma = Sma::new("sma1", 1).unwrap();
835        let expr = SignalExpr::raw("sma1").normalize(NormMethod::ZScore, 5);
836        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma)];
837        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
838
839        // Feed 5 bars of [10, 10, 10, 10, 10] → z-score of 10 is 0
840        for _ in 0..4 {
841            composed.update(&bar("10")).unwrap();
842        }
843        let v = composed.update(&bar("10")).unwrap();
844        if let SignalValue::Scalar(z) = v {
845            assert!(z.abs() < dec!(0.001), "z-score of mean should be near 0, got {z}");
846        } else {
847            panic!("expected Scalar");
848        }
849    }
850
851    // ── Normalize: Percentile ────────────────────────────────────────────────
852
853    #[test]
854    fn test_normalize_percentile_highest_value() {
855        let sma = Sma::new("sma1", 1).unwrap();
856        let expr = SignalExpr::raw("sma1").normalize(NormMethod::Percentile, 4);
857        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma)];
858        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
859
860        composed.update(&bar("10")).unwrap();
861        composed.update(&bar("20")).unwrap();
862        composed.update(&bar("30")).unwrap();
863        let v = composed.update(&bar("100")).unwrap(); // clearly the max
864        if let SignalValue::Scalar(pct) = v {
865            assert!(pct > dec!(0.5), "max value should have pct > 0.5, got {pct}");
866        } else {
867            panic!("expected Scalar");
868        }
869    }
870
871    // ── Threshold ────────────────────────────────────────────────────────────
872
873    #[test]
874    fn test_threshold_above_emits_one_when_above() {
875        let sma = Sma::new("sma1", 1).unwrap();
876        let expr = SignalExpr::raw("sma1").threshold(dec!(50), Direction::Above);
877        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma)];
878        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
879
880        let v = composed.update(&bar("75")).unwrap();
881        assert_eq!(v, SignalValue::Scalar(dec!(1)));
882    }
883
884    #[test]
885    fn test_threshold_above_emits_zero_when_below() {
886        let sma = Sma::new("sma1", 1).unwrap();
887        let expr = SignalExpr::raw("sma1").threshold(dec!(50), Direction::Above);
888        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma)];
889        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
890
891        let v = composed.update(&bar("25")).unwrap();
892        assert_eq!(v, SignalValue::Scalar(dec!(0)));
893    }
894
895    #[test]
896    fn test_threshold_below_emits_neg_one_when_below() {
897        let sma = Sma::new("sma1", 1).unwrap();
898        let expr = SignalExpr::raw("sma1").threshold(dec!(50), Direction::Below);
899        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma)];
900        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
901
902        let v = composed.update(&bar("20")).unwrap();
903        assert_eq!(v, SignalValue::Scalar(dec!(-1)));
904    }
905
906    #[test]
907    fn test_threshold_cross_emits_one_on_upward_cross() {
908        let sma = Sma::new("sma1", 1).unwrap();
909        let expr = SignalExpr::raw("sma1").threshold(dec!(50), Direction::Cross);
910        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma)];
911        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
912
913        composed.update(&bar("40")).unwrap(); // below threshold, prev = Unavailable
914        composed.update(&bar("40")).unwrap(); // prev = 40 (below)
915        let v = composed.update(&bar("60")).unwrap(); // crosses above
916        assert_eq!(v, SignalValue::Scalar(dec!(1)));
917    }
918
919    #[test]
920    fn test_threshold_cross_emits_neg_one_on_downward_cross() {
921        let sma = Sma::new("sma1", 1).unwrap();
922        let expr = SignalExpr::raw("sma1").threshold(dec!(50), Direction::Cross);
923        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma)];
924        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
925
926        composed.update(&bar("60")).unwrap(); // above
927        composed.update(&bar("60")).unwrap(); // prev = 60 (above)
928        let v = composed.update(&bar("40")).unwrap(); // crosses below
929        assert_eq!(v, SignalValue::Scalar(dec!(-1)));
930    }
931
932    // ── SignalBuilder ────────────────────────────────────────────────────────
933
934    #[test]
935    fn test_builder_builds_composed_signal() {
936        let sma = Sma::new("sma5", 5).unwrap();
937        let mut composed = SignalBuilder::new(sma).lag(1).build();
938        assert_eq!(composed.name(), "composed(sma5)");
939        let v = composed.update(&bar("100")).unwrap();
940        assert_eq!(v, SignalValue::Unavailable); // lag-1 not filled yet
941    }
942
943    #[test]
944    fn test_builder_build_named() {
945        let sma = Sma::new("sma5", 5).unwrap();
946        let composed = SignalBuilder::new(sma).build_named("my_signal");
947        assert_eq!(composed.name(), "my_signal");
948    }
949
950    #[test]
951    fn test_builder_scale() {
952        let sma = Sma::new("sma1", 1).unwrap();
953        let mut composed = SignalBuilder::new(sma).scale(dec!(3)).build();
954        let v = composed.update(&bar("10")).unwrap();
955        assert_eq!(v, SignalValue::Scalar(dec!(30)));
956    }
957
958    #[test]
959    fn test_builder_normalize_minmax() {
960        let sma = Sma::new("sma1", 1).unwrap();
961        let mut composed = SignalBuilder::new(sma)
962            .normalize_window(NormMethod::MinMax, 3)
963            .build();
964        composed.update(&bar("0")).unwrap();
965        composed.update(&bar("50")).unwrap();
966        let v = composed.update(&bar("100")).unwrap();
967        assert_eq!(v, SignalValue::Scalar(dec!(1)));
968    }
969
970    #[test]
971    fn test_builder_threshold_above() {
972        let sma = Sma::new("sma1", 1).unwrap();
973        let mut composed = SignalBuilder::new(sma)
974            .threshold(dec!(50), Direction::Above)
975            .build();
976        let v = composed.update(&bar("80")).unwrap();
977        assert_eq!(v, SignalValue::Scalar(dec!(1)));
978    }
979
980    #[test]
981    fn test_builder_chain_lag_normalize_threshold() {
982        let rsi = Rsi::new("rsi5", 5).unwrap();
983        let mut composed = SignalBuilder::new(rsi)
984            .lag(1)
985            .normalize_window(NormMethod::ZScore, 10)
986            .threshold(dec!(1), Direction::Above)
987            .build();
988
989        // Feed enough bars to warm up RSI + lag + z-score window
990        for _ in 0..30 {
991            composed.update(&bar("50")).unwrap();
992        }
993        // All warming should be done; the value is Scalar(0) since all bars are flat (z=0)
994        let v = composed.update(&bar("50")).unwrap();
995        assert!(matches!(v, SignalValue::Scalar(_)));
996    }
997
998    // ── ComposedSignal reset ─────────────────────────────────────────────────
999
1000    #[test]
1001    fn test_composed_reset_restarts_warmup() {
1002        let sma = Sma::new("sma3", 3).unwrap();
1003        let expr = SignalExpr::raw("sma3");
1004        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma)];
1005        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
1006
1007        feed_n(&mut composed, "10", 3);
1008        assert!(composed.is_ready());
1009        composed.reset();
1010        assert!(!composed.is_ready());
1011        let v = composed.update(&bar("10")).unwrap();
1012        assert_eq!(v, SignalValue::Unavailable);
1013    }
1014
1015    // ── Period reporting ─────────────────────────────────────────────────────
1016
1017    #[test]
1018    fn test_composed_period_reflects_max_leaf_period() {
1019        let ema = Ema::new("ema10", 10).unwrap();
1020        let composed = SignalBuilder::new(ema).build();
1021        assert_eq!(composed.period(), 10);
1022    }
1023
1024    // ── Error cases ──────────────────────────────────────────────────────────
1025
1026    #[test]
1027    fn test_composed_new_fails_with_empty_leaves() {
1028        let expr = SignalExpr::raw("nonexistent");
1029        let leaves: Vec<Box<dyn Signal>> = vec![];
1030        let result = ComposedSignal::new("test", expr, leaves);
1031        assert!(result.is_err());
1032    }
1033}