qs-strategy 0.4.2

Synchronous reusable configured strategy core
Documentation
use crate::material::{
    MaterialBuild, MaterialEvalContext, MaterialEvaluator, MaterialFactory, MaterialLookback,
    MaterialUpdateTrigger, ParamKind, ParamSpec,
};
use crate::{
    MaterialArg, MaterialArgs, NormalizedCalculation, NumericCalculation, NumericDescriptor,
    NumericInputs, NumericMissingPolicy, NumericRange, NumericUnit, ScalarType, SourceId, Value,
    ValueType,
};
use std::sync::Arc;
pub const MATERIAL_MACD_ATR: &str = "macd_atr";
pub const MATERIAL_MACD_SIGNAL_ATR: &str = "macd_signal_atr";
pub const MATERIAL_MACD_HISTOGRAM_ATR: &str = "macd_histogram_atr";
pub const MATERIAL_MA_DISTANCE_ATR: &str = "ma_distance_atr";
pub const MATERIAL_MA_SLOPE_ATR: &str = "ma_slope_atr";
pub const MATERIAL_MA_ACCELERATION_ATR: &str = "ma_acceleration_atr";
pub const MATERIAL_MA_GAP_ATR: &str = "ma_gap_atr";
const KEYS: &[&str] = &[
    MATERIAL_MACD_ATR,
    MATERIAL_MACD_SIGNAL_ATR,
    MATERIAL_MACD_HISTOGRAM_ATR,
    MATERIAL_MA_DISTANCE_ATR,
    MATERIAL_MA_SLOPE_ATR,
    MATERIAL_MA_ACCELERATION_ATR,
    MATERIAL_MA_GAP_ATR,
];
const SCHEMA: [ParamSpec; 1] = [ParamSpec {
    name: "source",
    kind: ParamKind::Source,
    required: true,
}];
pub(crate) fn registrations() -> impl Iterator<Item = (&'static str, Arc<dyn MaterialFactory>)> {
    KEYS.iter().copied().map(|key| {
        (
            key,
            Arc::new(NormalizationFactory { key }) as Arc<dyn MaterialFactory>,
        )
    })
}
struct NormalizationFactory {
    key: &'static str,
}
impl NormalizationFactory {
    fn source(p: &MaterialArgs) -> Result<SourceId, String> {
        match p.get("source") {
            Some(MaterialArg::Source(v)) => Ok(v.clone()),
            _ => Err("normalization requires source".into()),
        }
    }
    fn calculation(&self) -> NormalizedCalculation {
        match self.key {
            MATERIAL_MACD_ATR => NormalizedCalculation::MacdAtr,
            MATERIAL_MACD_SIGNAL_ATR => NormalizedCalculation::MacdSignalAtr,
            MATERIAL_MACD_HISTOGRAM_ATR => NormalizedCalculation::MacdHistogramAtr,
            MATERIAL_MA_DISTANCE_ATR => NormalizedCalculation::MaDistanceAtr,
            MATERIAL_MA_SLOPE_ATR => NormalizedCalculation::MaSlopeAtr,
            MATERIAL_MA_ACCELERATION_ATR => NormalizedCalculation::MaAccelerationAtr,
            _ => NormalizedCalculation::MaGapAtr,
        }
    }
}
impl MaterialFactory for NormalizationFactory {
    fn numeric_descriptor(
        &self,
        p: &MaterialArgs,
        inputs: &[ValueType],
    ) -> Result<Option<NumericDescriptor>, String> {
        let numerator = if self.key == MATERIAL_MA_SLOPE_ATR {
            ScalarType::PricePerObservation
        } else if self.key == MATERIAL_MA_ACCELERATION_ATR {
            ScalarType::PricePerObservationSquared
        } else {
            ScalarType::Price
        };
        if inputs.len() != 2
            || inputs[0].scalar != numerator
            || inputs[1].scalar != ScalarType::Price
        {
            return Err("ATR normalization requires a typed numerator and Price ATR".into());
        }
        let (output, unit) = if numerator == ScalarType::PricePerObservation {
            (
                ScalarType::RatioPerObservation,
                NumericUnit::RatioPerObservation,
            )
        } else if numerator == ScalarType::PricePerObservationSquared {
            (
                ScalarType::RatioPerObservationSquared,
                NumericUnit::RatioPerObservationSquared,
            )
        } else {
            (ScalarType::Ratio, NumericUnit::Ratio)
        };
        Ok(Some(NumericDescriptor {
            calculation: NumericCalculation::NormalizedPair {
                calculation: self.calculation(),
            },
            source_clock: Self::source(p)?,
            inputs: NumericInputs::Scalar(inputs.to_vec()),
            output_type: ValueType::optional(output),
            unit,
            range: NumericRange::Unbounded,
            missing: NumericMissingPolicy::ConsumeWindowSlot,
            first_output_observations: 1,
            required_lookback: 1,
            max_state_bytes: std::mem::size_of::<NormalizationEvaluator>(),
            exact_aliases: &[],
        }))
    }
    fn params(&self) -> &[ParamSpec] {
        &SCHEMA
    }
    fn build(&self, p: &MaterialArgs, inputs: &[ValueType]) -> Result<MaterialBuild, String> {
        let d = self.numeric_descriptor(p, inputs)?.unwrap();
        Ok(MaterialBuild {
            output_type: d.output_type,
            lookback: MaterialLookback::InheritInputs { minimum: 1 },
            max_state_bytes: d.max_state_bytes,
            evaluator: Box::new(NormalizationEvaluator {
                output: d.output_type.scalar,
            }),
        })
    }
    fn update_trigger(
        &self,
        p: &MaterialArgs,
        _: &[ValueType],
    ) -> Result<MaterialUpdateTrigger, String> {
        Ok(MaterialUpdateTrigger::Source(Self::source(p)?))
    }
}
#[derive(Clone)]
struct NormalizationEvaluator {
    output: ScalarType,
}
impl MaterialEvaluator for NormalizationEvaluator {
    fn clone_box(&self) -> Box<dyn MaterialEvaluator> {
        Box::new(self.clone())
    }
    fn evaluate(
        &mut self,
        inputs: &[Value],
        context: &MaterialEvalContext<'_>,
    ) -> Result<Value, String> {
        if context.input_updates.iter().any(|updated| !*updated) {
            return Ok(Value::Missing(self.output));
        }
        let numerator = number(&inputs[0])?;
        let denominator = number(&inputs[1])?;
        let (Some(n), Some(d)) = (numerator, denominator) else {
            return Ok(Value::Missing(self.output));
        };
        if d <= 0.0 {
            return Ok(Value::Missing(self.output));
        }
        Ok(match self.output {
            ScalarType::Ratio => Value::Ratio(n / d),
            ScalarType::RatioPerObservation => Value::RatioPerObservation(n / d),
            ScalarType::RatioPerObservationSquared => Value::RatioPerObservationSquared(n / d),
            _ => unreachable!(),
        })
    }
}
fn number(v: &Value) -> Result<Option<f64>, String> {
    match v {
        Value::Missing(_) => Ok(None),
        Value::Price(v) | Value::PricePerObservation(v) | Value::PricePerObservationSquared(v)
            if v.is_finite() =>
        {
            Ok(Some(*v))
        }
        _ => Err("normalization input is not a finite compatible scalar".into()),
    }
}
#[cfg(test)]
mod tests {
    use super::*;
    use crate::StrategyInput;
    use chrono::NaiveDateTime;
    #[test]
    fn normalized_pairs_preserve_units_and_zero_domain() {
        let source = SourceId::new("bars").unwrap();
        let input = StrategyInput {
            time: NaiveDateTime::default(),
            ready: true,
            completed_bars: vec![],
            values: vec![],
            trade_slots: vec![],
            feedback: vec![],
        };
        for (key, types, values, expected) in [
            (
                MATERIAL_MACD_ATR,
                [
                    ValueType::optional(ScalarType::Price),
                    ValueType::optional(ScalarType::Price),
                ],
                [Value::Price(2.0), Value::Price(4.0)],
                Value::Ratio(0.5),
            ),
            (
                MATERIAL_MA_SLOPE_ATR,
                [
                    ValueType::optional(ScalarType::PricePerObservation),
                    ValueType::optional(ScalarType::Price),
                ],
                [Value::PricePerObservation(2.0), Value::Price(4.0)],
                Value::RatioPerObservation(0.5),
            ),
            (
                MATERIAL_MA_ACCELERATION_ATR,
                [
                    ValueType::optional(ScalarType::PricePerObservationSquared),
                    ValueType::optional(ScalarType::Price),
                ],
                [Value::PricePerObservationSquared(2.0), Value::Price(4.0)],
                Value::RatioPerObservationSquared(0.5),
            ),
        ] {
            let p = MaterialArgs::new([("source", MaterialArg::Source(source.clone()))]);
            let mut e = NormalizationFactory { key }
                .build(&p, &types)
                .unwrap()
                .evaluator;
            let context = MaterialEvalContext {
                input: &input,
                input_updates: &[true, true],
                any_input_updates: &[true, true],
                feedback: &[],
                retained_feedback: &[],
            };
            assert_eq!(e.evaluate(&values, &context).unwrap(), expected);
            let zero = [values[0].clone(), Value::Price(0.0)];
            assert_eq!(
                e.evaluate(&zero, &context).unwrap(),
                Value::Missing(expected.scalar_type())
            );
        }
    }
}