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                        let result = match (&v, &*prev) {
371                            (SignalValue::Scalar(curr), SignalValue::Scalar(p)) => {
372                                if *curr > *level && *p <= *level {
373                                    SignalValue::Scalar(Decimal::ONE)
374                                } else if *curr < *level && *p >= *level {
375                                    SignalValue::Scalar(-Decimal::ONE)
376                                } else {
377                                    SignalValue::Scalar(Decimal::ZERO)
378                                }
379                            }
380                            _ => SignalValue::Unavailable,
381                        };
382                        result
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    ///
662    /// # Panics
663    ///
664    /// Panics if signal construction fails (which cannot happen here since the
665    /// leaf is already valid).
666    pub fn build(self) -> ComposedSignal {
667        let leaf_name = self.signal.name().to_owned();
668        let composed_name = format!("composed({})", leaf_name);
669        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(self.signal)];
670        // Safety: leaves is non-empty by construction.
671        ComposedSignal::new(composed_name, self.expr, leaves)
672            .expect("ComposedSignal construction with a valid leaf signal cannot fail")
673    }
674
675    /// Consumes the builder and produces a [`ComposedSignal`] with an explicit name.
676    pub fn build_named(self, name: impl Into<String>) -> ComposedSignal {
677        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(self.signal)];
678        ComposedSignal::new(name, self.expr, leaves)
679            .expect("ComposedSignal construction with a valid leaf signal cannot fail")
680    }
681}
682
683// ── Tests ─────────────────────────────────────────────────────────────────────
684
685#[cfg(test)]
686mod tests {
687    use super::*;
688    use crate::signals::indicators::{Ema, Rsi, Sma};
689    use rust_decimal_macros::dec;
690
691    fn bar(close: &str) -> BarInput {
692        BarInput::from_close(close.parse().unwrap())
693    }
694
695    fn feed_n(signal: &mut impl Signal, close: &str, n: usize) {
696        for _ in 0..n {
697            signal.update(&bar(close)).unwrap();
698        }
699    }
700
701    // ── SignalExpr leaf_names ────────────────────────────────────────────────
702
703    #[test]
704    fn test_expr_raw_leaf_name() {
705        let expr = SignalExpr::raw("sma5");
706        assert_eq!(expr.leaf_names(), vec!["sma5"]);
707    }
708
709    #[test]
710    fn test_expr_add_leaf_names() {
711        let expr = SignalExpr::raw("a").add(SignalExpr::raw("b"));
712        let names = expr.leaf_names();
713        assert!(names.contains(&"a"));
714        assert!(names.contains(&"b"));
715    }
716
717    #[test]
718    fn test_expr_nested_leaf_names() {
719        let expr = SignalExpr::raw("sma5")
720            .lag(1)
721            .normalize(NormMethod::ZScore, 20)
722            .threshold(dec!(0), Direction::Above);
723        assert_eq!(expr.leaf_names(), vec!["sma5"]);
724    }
725
726    // ── ComposedSignal: Raw passthrough ──────────────────────────────────────
727
728    #[test]
729    fn test_composed_raw_passthrough() {
730        let sma = Sma::new("sma3", 3).unwrap();
731        let expr = SignalExpr::raw("sma3");
732        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma)];
733        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
734
735        feed_n(&mut composed, "10", 2);
736        let v = composed.update(&bar("10")).unwrap();
737        assert!(matches!(v, SignalValue::Scalar(_)));
738    }
739
740    #[test]
741    fn test_composed_raw_unavailable_during_warmup() {
742        let sma = Sma::new("sma5", 5).unwrap();
743        let expr = SignalExpr::raw("sma5");
744        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma)];
745        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
746
747        let v = composed.update(&bar("10")).unwrap();
748        assert_eq!(v, SignalValue::Unavailable);
749    }
750
751    // ── Mul ──────────────────────────────────────────────────────────────────
752
753    #[test]
754    fn test_composed_mul_scales_value() {
755        let sma = Sma::new("sma1", 1).unwrap();
756        let expr = SignalExpr::raw("sma1").mul(dec!(2));
757        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma)];
758        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
759
760        let v = composed.update(&bar("10")).unwrap();
761        assert_eq!(v, SignalValue::Scalar(dec!(20)));
762    }
763
764    // ── Add ──────────────────────────────────────────────────────────────────
765
766    #[test]
767    fn test_composed_add_two_signals() {
768        let sma1 = Sma::new("sma_a", 1).unwrap();
769        let sma2 = Sma::new("sma_b", 1).unwrap();
770        let expr = SignalExpr::raw("sma_a").add(SignalExpr::raw("sma_b"));
771        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma1), Box::new(sma2)];
772        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
773
774        let v = composed.update(&bar("15")).unwrap();
775        // sma_a = 15, sma_b = 15, sum = 30
776        assert_eq!(v, SignalValue::Scalar(dec!(30)));
777    }
778
779    // ── Lag ──────────────────────────────────────────────────────────────────
780
781    #[test]
782    fn test_composed_lag_delays_values() {
783        let sma = Sma::new("sma1", 1).unwrap();
784        let expr = SignalExpr::raw("sma1").lag(2);
785        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma)];
786        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
787
788        // bars 1,2: buffer fills but has not produced 2 prior values yet
789        composed.update(&bar("10")).unwrap();
790        composed.update(&bar("20")).unwrap();
791        // bar 3: lag-2 should produce bar-1's value = 10
792        let v = composed.update(&bar("30")).unwrap();
793        assert_eq!(v, SignalValue::Scalar(dec!(10)));
794    }
795
796    #[test]
797    fn test_composed_lag_zero_is_passthrough() {
798        let sma = Sma::new("sma1", 1).unwrap();
799        let expr = SignalExpr::raw("sma1").lag(0);
800        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma)];
801        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
802
803        let v = composed.update(&bar("42")).unwrap();
804        assert_eq!(v, SignalValue::Scalar(dec!(42)));
805    }
806
807    // ── Normalize: MinMax ────────────────────────────────────────────────────
808
809    #[test]
810    fn test_normalize_minmax_range_of_constant_returns_zero() {
811        let sma = Sma::new("sma1", 1).unwrap();
812        let expr = SignalExpr::raw("sma1").normalize(NormMethod::MinMax, 3);
813        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma)];
814        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
815
816        // Feed 3 identical values → min==max → output 0
817        composed.update(&bar("10")).unwrap();
818        composed.update(&bar("10")).unwrap();
819        let v = composed.update(&bar("10")).unwrap();
820        assert_eq!(v, SignalValue::Scalar(dec!(0)));
821    }
822
823    #[test]
824    fn test_normalize_minmax_high_value_approaches_one() {
825        let sma = Sma::new("sma1", 1).unwrap();
826        let expr = SignalExpr::raw("sma1").normalize(NormMethod::MinMax, 3);
827        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma)];
828        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
829
830        // window: [0, 50, 100] → current=100, min=0, max=100 → 1.0
831        composed.update(&bar("0")).unwrap();
832        composed.update(&bar("50")).unwrap();
833        let v = composed.update(&bar("100")).unwrap();
834        assert_eq!(v, SignalValue::Scalar(dec!(1)));
835    }
836
837    // ── Normalize: ZScore ────────────────────────────────────────────────────
838
839    #[test]
840    fn test_normalize_zscore_mean_value_near_zero() {
841        let sma = Sma::new("sma1", 1).unwrap();
842        let expr = SignalExpr::raw("sma1").normalize(NormMethod::ZScore, 5);
843        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma)];
844        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
845
846        // Feed 5 bars of [10, 10, 10, 10, 10] → z-score of 10 is 0
847        for _ in 0..4 {
848            composed.update(&bar("10")).unwrap();
849        }
850        let v = composed.update(&bar("10")).unwrap();
851        if let SignalValue::Scalar(z) = v {
852            assert!(z.abs() < dec!(0.001), "z-score of mean should be near 0, got {z}");
853        } else {
854            panic!("expected Scalar");
855        }
856    }
857
858    // ── Normalize: Percentile ────────────────────────────────────────────────
859
860    #[test]
861    fn test_normalize_percentile_highest_value() {
862        let sma = Sma::new("sma1", 1).unwrap();
863        let expr = SignalExpr::raw("sma1").normalize(NormMethod::Percentile, 4);
864        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma)];
865        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
866
867        composed.update(&bar("10")).unwrap();
868        composed.update(&bar("20")).unwrap();
869        composed.update(&bar("30")).unwrap();
870        let v = composed.update(&bar("100")).unwrap(); // clearly the max
871        if let SignalValue::Scalar(pct) = v {
872            assert!(pct > dec!(0.5), "max value should have pct > 0.5, got {pct}");
873        } else {
874            panic!("expected Scalar");
875        }
876    }
877
878    // ── Threshold ────────────────────────────────────────────────────────────
879
880    #[test]
881    fn test_threshold_above_emits_one_when_above() {
882        let sma = Sma::new("sma1", 1).unwrap();
883        let expr = SignalExpr::raw("sma1").threshold(dec!(50), Direction::Above);
884        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma)];
885        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
886
887        let v = composed.update(&bar("75")).unwrap();
888        assert_eq!(v, SignalValue::Scalar(dec!(1)));
889    }
890
891    #[test]
892    fn test_threshold_above_emits_zero_when_below() {
893        let sma = Sma::new("sma1", 1).unwrap();
894        let expr = SignalExpr::raw("sma1").threshold(dec!(50), Direction::Above);
895        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma)];
896        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
897
898        let v = composed.update(&bar("25")).unwrap();
899        assert_eq!(v, SignalValue::Scalar(dec!(0)));
900    }
901
902    #[test]
903    fn test_threshold_below_emits_neg_one_when_below() {
904        let sma = Sma::new("sma1", 1).unwrap();
905        let expr = SignalExpr::raw("sma1").threshold(dec!(50), Direction::Below);
906        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma)];
907        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
908
909        let v = composed.update(&bar("20")).unwrap();
910        assert_eq!(v, SignalValue::Scalar(dec!(-1)));
911    }
912
913    #[test]
914    fn test_threshold_cross_emits_one_on_upward_cross() {
915        let sma = Sma::new("sma1", 1).unwrap();
916        let expr = SignalExpr::raw("sma1").threshold(dec!(50), Direction::Cross);
917        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma)];
918        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
919
920        composed.update(&bar("40")).unwrap(); // below threshold, prev = Unavailable
921        composed.update(&bar("40")).unwrap(); // prev = 40 (below)
922        let v = composed.update(&bar("60")).unwrap(); // crosses above
923        assert_eq!(v, SignalValue::Scalar(dec!(1)));
924    }
925
926    #[test]
927    fn test_threshold_cross_emits_neg_one_on_downward_cross() {
928        let sma = Sma::new("sma1", 1).unwrap();
929        let expr = SignalExpr::raw("sma1").threshold(dec!(50), Direction::Cross);
930        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma)];
931        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
932
933        composed.update(&bar("60")).unwrap(); // above
934        composed.update(&bar("60")).unwrap(); // prev = 60 (above)
935        let v = composed.update(&bar("40")).unwrap(); // crosses below
936        assert_eq!(v, SignalValue::Scalar(dec!(-1)));
937    }
938
939    // ── SignalBuilder ────────────────────────────────────────────────────────
940
941    #[test]
942    fn test_builder_builds_composed_signal() {
943        let sma = Sma::new("sma5", 5).unwrap();
944        let mut composed = SignalBuilder::new(sma).lag(1).build();
945        assert_eq!(composed.name(), "composed(sma5)");
946        let v = composed.update(&bar("100")).unwrap();
947        assert_eq!(v, SignalValue::Unavailable); // lag-1 not filled yet
948    }
949
950    #[test]
951    fn test_builder_build_named() {
952        let sma = Sma::new("sma5", 5).unwrap();
953        let composed = SignalBuilder::new(sma).build_named("my_signal");
954        assert_eq!(composed.name(), "my_signal");
955    }
956
957    #[test]
958    fn test_builder_scale() {
959        let sma = Sma::new("sma1", 1).unwrap();
960        let mut composed = SignalBuilder::new(sma).scale(dec!(3)).build();
961        let v = composed.update(&bar("10")).unwrap();
962        assert_eq!(v, SignalValue::Scalar(dec!(30)));
963    }
964
965    #[test]
966    fn test_builder_normalize_minmax() {
967        let sma = Sma::new("sma1", 1).unwrap();
968        let mut composed = SignalBuilder::new(sma)
969            .normalize_window(NormMethod::MinMax, 3)
970            .build();
971        composed.update(&bar("0")).unwrap();
972        composed.update(&bar("50")).unwrap();
973        let v = composed.update(&bar("100")).unwrap();
974        assert_eq!(v, SignalValue::Scalar(dec!(1)));
975    }
976
977    #[test]
978    fn test_builder_threshold_above() {
979        let sma = Sma::new("sma1", 1).unwrap();
980        let mut composed = SignalBuilder::new(sma)
981            .threshold(dec!(50), Direction::Above)
982            .build();
983        let v = composed.update(&bar("80")).unwrap();
984        assert_eq!(v, SignalValue::Scalar(dec!(1)));
985    }
986
987    #[test]
988    fn test_builder_chain_lag_normalize_threshold() {
989        let rsi = Rsi::new("rsi5", 5).unwrap();
990        let mut composed = SignalBuilder::new(rsi)
991            .lag(1)
992            .normalize_window(NormMethod::ZScore, 10)
993            .threshold(dec!(1), Direction::Above)
994            .build();
995
996        // Feed enough bars to warm up RSI + lag + z-score window
997        for _ in 0..30 {
998            composed.update(&bar("50")).unwrap();
999        }
1000        // All warming should be done; the value is Scalar(0) since all bars are flat (z=0)
1001        let v = composed.update(&bar("50")).unwrap();
1002        assert!(matches!(v, SignalValue::Scalar(_)));
1003    }
1004
1005    // ── ComposedSignal reset ─────────────────────────────────────────────────
1006
1007    #[test]
1008    fn test_composed_reset_restarts_warmup() {
1009        let sma = Sma::new("sma3", 3).unwrap();
1010        let expr = SignalExpr::raw("sma3");
1011        let leaves: Vec<Box<dyn Signal>> = vec![Box::new(sma)];
1012        let mut composed = ComposedSignal::new("test", expr, leaves).unwrap();
1013
1014        feed_n(&mut composed, "10", 3);
1015        assert!(composed.is_ready());
1016        composed.reset();
1017        assert!(!composed.is_ready());
1018        let v = composed.update(&bar("10")).unwrap();
1019        assert_eq!(v, SignalValue::Unavailable);
1020    }
1021
1022    // ── Period reporting ─────────────────────────────────────────────────────
1023
1024    #[test]
1025    fn test_composed_period_reflects_max_leaf_period() {
1026        let ema = Ema::new("ema10", 10).unwrap();
1027        let composed = SignalBuilder::new(ema).build();
1028        assert_eq!(composed.period(), 10);
1029    }
1030
1031    // ── Error cases ──────────────────────────────────────────────────────────
1032
1033    #[test]
1034    fn test_composed_new_fails_with_empty_leaves() {
1035        let expr = SignalExpr::raw("nonexistent");
1036        let leaves: Vec<Box<dyn Signal>> = vec![];
1037        let result = ComposedSignal::new("test", expr, leaves);
1038        assert!(result.is_err());
1039    }
1040}