use crate::error::FinError;
use crate::signals::{BarInput, Signal, SignalValue};
use rust_decimal::Decimal;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CompositeMode {
WeightedSum,
All,
Any,
First,
}
struct Constituent {
signal: Box<dyn Signal + Send>,
weight: Decimal,
}
pub struct CompositeSignal {
name: String,
constituents: Vec<Constituent>,
mode: CompositeMode,
}
impl std::fmt::Debug for CompositeSignal {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("CompositeSignal")
.field("name", &self.name)
.field("mode", &self.mode)
.field("constituents", &self.constituents.len())
.finish()
}
}
pub struct CompositeBuilder {
name: String,
constituents: Vec<Constituent>,
mode: CompositeMode,
}
impl CompositeSignal {
pub fn builder(name: impl Into<String>) -> CompositeBuilder {
CompositeBuilder {
name: name.into(),
constituents: Vec::new(),
mode: CompositeMode::WeightedSum,
}
}
}
impl CompositeBuilder {
#[must_use]
pub fn add(mut self, signal: impl Signal + Send + 'static, weight: Decimal) -> Self {
self.constituents.push(Constituent {
signal: Box::new(signal),
weight,
});
self
}
#[must_use]
pub fn mode(mut self, mode: CompositeMode) -> Self {
self.mode = mode;
self
}
pub fn build(self) -> CompositeSignal {
assert!(
!self.constituents.is_empty(),
"CompositeSignal '{}' must have at least one constituent",
self.name
);
CompositeSignal {
name: self.name,
constituents: self.constituents,
mode: self.mode,
}
}
}
impl Signal for CompositeSignal {
fn name(&self) -> &str {
&self.name
}
fn is_ready(&self) -> bool {
self.constituents.iter().all(|c| c.signal.is_ready())
}
fn period(&self) -> usize {
self.constituents.iter().map(|c| c.signal.period()).max().unwrap_or(0)
}
fn reset(&mut self) {
for c in &mut self.constituents {
c.signal.reset();
}
}
fn update(&mut self, bar: &BarInput) -> Result<SignalValue, FinError> {
let mut values: Vec<(Decimal, Decimal)> = Vec::with_capacity(self.constituents.len()); let mut any_unavailable = false;
let mut all_unavailable = true;
for c in &mut self.constituents {
match c.signal.update(bar)? {
SignalValue::Scalar(v) => {
values.push((c.weight, v));
all_unavailable = false;
}
SignalValue::Unavailable => {
any_unavailable = true;
values.push((c.weight, Decimal::ZERO)); }
}
}
match self.mode {
CompositeMode::WeightedSum => {
if any_unavailable {
return Ok(SignalValue::Unavailable);
}
let total_weight: Decimal = values.iter().map(|(w, _)| *w).sum();
if total_weight == Decimal::ZERO {
return Ok(SignalValue::Unavailable);
}
let weighted_sum: Decimal = values.iter().map(|(w, v)| *w * *v).sum();
let result = weighted_sum
.checked_div(total_weight)
.ok_or(FinError::ArithmeticOverflow)?;
Ok(SignalValue::Scalar(result))
}
CompositeMode::All => {
if any_unavailable {
return Ok(SignalValue::Unavailable);
}
let all_nonzero = values.iter().all(|(_, v)| !v.is_zero());
Ok(SignalValue::Scalar(if all_nonzero {
Decimal::ONE
} else {
Decimal::ZERO
}))
}
CompositeMode::Any => {
if all_unavailable {
return Ok(SignalValue::Unavailable);
}
let any_nonzero = values
.iter()
.enumerate()
.filter(|(i, _)| {
let _ = i; true
})
.any(|(_, (_, v))| !v.is_zero());
let any_nonzero = self.any_nonzero_available(&values, any_unavailable);
Ok(SignalValue::Scalar(if any_nonzero {
Decimal::ONE
} else {
Decimal::ZERO
}))
}
CompositeMode::First => {
for (i, c) in self.constituents.iter_mut().enumerate() {
let _ = c; if let Some((_, v)) = values.get(i) {
let _ = v;
}
}
self.first_available(&values, any_unavailable)
}
}
}
}
impl CompositeSignal {
fn any_nonzero_available(&self, values: &[(Decimal, Decimal)], any_unavailable: bool) -> bool {
if !any_unavailable {
return values.iter().any(|(_, v)| !v.is_zero());
}
values.iter().any(|(_, v)| !v.is_zero())
}
fn first_available(
&self,
values: &[(Decimal, Decimal)],
any_unavailable: bool,
) -> Result<SignalValue, FinError> {
if !any_unavailable {
if let Some((_, v)) = values.first() {
return Ok(SignalValue::Scalar(*v));
}
}
for (_, v) in values {
if !v.is_zero() {
return Ok(SignalValue::Scalar(*v));
}
}
Ok(SignalValue::Unavailable)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::signals::indicators::Sma;
use rust_decimal_macros::dec;
fn bar(close: Decimal) -> BarInput {
BarInput::from_close(close)
}
fn warmed_composite(mode: CompositeMode) -> CompositeSignal {
CompositeSignal::builder("test")
.add(Sma::new("sma2", 2).unwrap(), dec!(1))
.add(Sma::new("sma2b", 2).unwrap(), dec!(1))
.mode(mode)
.build()
}
fn warm_up(sig: &mut CompositeSignal, n: usize) {
for _ in 0..n {
let _ = sig.update(&bar(dec!(10)));
}
}
#[test]
fn weighted_sum_unavailable_before_warmup() {
let mut sig = warmed_composite(CompositeMode::WeightedSum);
assert_eq!(sig.update(&bar(dec!(10))).unwrap(), SignalValue::Unavailable);
}
#[test]
fn weighted_sum_available_after_warmup() {
let mut sig = warmed_composite(CompositeMode::WeightedSum);
warm_up(&mut sig, 2);
let v = sig.update(&bar(dec!(10))).unwrap();
assert!(matches!(v, SignalValue::Scalar(_)));
}
#[test]
fn weighted_sum_equal_weights_is_average() {
let mut sig = CompositeSignal::builder("avg")
.add(Sma::new("sma2a", 2).unwrap(), dec!(1))
.add(Sma::new("sma2b", 2).unwrap(), dec!(1))
.mode(CompositeMode::WeightedSum)
.build();
let _ = sig.update(&bar(dec!(10))).unwrap();
let v = sig.update(&bar(dec!(20))).unwrap();
if let SignalValue::Scalar(val) = v {
assert_eq!(val, dec!(15), "expected (15+15)/2 = 15, got {val}");
} else {
panic!("expected Scalar");
}
}
#[test]
fn all_mode_requires_all_nonzero() {
let mut sig = warmed_composite(CompositeMode::All);
warm_up(&mut sig, 2);
let v = sig.update(&bar(dec!(10))).unwrap();
assert_eq!(v, SignalValue::Scalar(dec!(1)));
}
#[test]
fn any_mode_available_after_partial_warmup() {
let mut sig = CompositeSignal::builder("any_test")
.add(Sma::new("sma1", 1).unwrap(), dec!(1))
.mode(CompositeMode::Any)
.build();
let v = sig.update(&bar(dec!(10))).unwrap();
assert_eq!(v, SignalValue::Scalar(dec!(1)));
}
#[test]
#[should_panic(expected = "must have at least one constituent")]
fn builder_panics_with_no_constituents() {
CompositeSignal::builder("empty").build();
}
}