use std::collections::HashMap;
#[cfg(feature = "serde")]
use serde::{Deserialize, Serialize};
use crate::SignalDirection as Direction;
#[derive(Debug, Clone, PartialEq)]
pub struct DirectionalStatement<T> {
pub tf: T,
pub direction: Direction,
pub strength: f64,
}
#[derive(Debug, Clone, PartialEq)]
pub struct Agreement {
pub direction: Direction,
pub agreement: f64,
pub conflict: bool,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[cfg_attr(
feature = "serde",
derive(Serialize, Deserialize),
serde(rename_all = "snake_case")
)]
pub enum AgreementStrategy {
#[default]
Majority,
WeightedByStrength,
TimeframeTopDown,
}
pub fn aggregate_agreement<T: Copy + Eq>(
mut statements: Vec<DirectionalStatement<T>>,
strategies: &[AgreementStrategy],
tf_order: &[T],
) -> Agreement {
let mut tally_by_strength = false;
for strategy in strategies {
match strategy {
AgreementStrategy::TimeframeTopDown => {
statements = apply_top_down(statements, tf_order);
}
AgreementStrategy::Majority => tally_by_strength = false,
AgreementStrategy::WeightedByStrength => tally_by_strength = true,
}
}
let result = tally(&statements, tally_by_strength);
if result.direction != Direction::Neutral {
return result;
}
let tie_broken = tally(
&apply_top_down(statements.clone(), tf_order),
tally_by_strength,
);
if tie_broken.direction != Direction::Neutral {
return tie_broken;
}
Agreement {
conflict: has_opposing_statements(&statements),
..result
}
}
fn has_opposing_statements<T>(statements: &[DirectionalStatement<T>]) -> bool {
let (mut bull, mut bear) = (false, false);
for s in statements {
match s.direction {
Direction::Bullish => bull = true,
Direction::Bearish => bear = true,
Direction::Neutral => {}
}
}
bull && bear
}
fn apply_top_down<T: Copy + Eq>(
statements: Vec<DirectionalStatement<T>>,
tf_order: &[T],
) -> Vec<DirectionalStatement<T>> {
let bias = tf_order.iter().find_map(|&tf| {
let on_tf: Vec<&DirectionalStatement<T>> =
statements.iter().filter(|s| s.tf == tf).collect();
if on_tf.is_empty() {
return None;
}
match tally_direction(on_tf.into_iter().cloned(), false) {
Direction::Neutral => None,
dir => Some(dir),
}
});
match bias {
None => statements,
Some(bias) => statements
.into_iter()
.filter(|s| s.direction == bias)
.collect(),
}
}
fn tally<T: Clone>(statements: &[DirectionalStatement<T>], by_strength: bool) -> Agreement {
let direction = tally_direction(statements.iter().cloned(), by_strength);
let (bull, bear): (f64, f64) = statements.iter().fold((0.0, 0.0), |(b, r), s| {
let weight = if by_strength { s.strength } else { 1.0 };
match s.direction {
Direction::Bullish => (b + weight, r),
Direction::Bearish => (b, r + weight),
Direction::Neutral => (b, r),
}
});
let total = bull + bear;
let agreement = if total > 0.0 {
(bull.max(bear)) / total
} else {
0.0
};
Agreement {
direction,
agreement,
conflict: false,
}
}
fn tally_direction<T>(
statements: impl Iterator<Item = DirectionalStatement<T>>,
by_strength: bool,
) -> Direction {
let mut weight: HashMap<Direction, f64> = HashMap::new();
for s in statements {
if s.direction == Direction::Neutral {
continue;
}
let w = if by_strength { s.strength } else { 1.0 };
*weight.entry(s.direction).or_insert(0.0) += w;
}
let bull = weight.get(&Direction::Bullish).copied().unwrap_or(0.0);
let bear = weight.get(&Direction::Bearish).copied().unwrap_or(0.0);
if bull == 0.0 && bear == 0.0 {
Direction::Neutral
} else if bull > bear {
Direction::Bullish
} else if bear > bull {
Direction::Bearish
} else {
Direction::Neutral }
}