kestrel_chartkit/scoring/
agreement.rs1use std::collections::HashMap;
2
3#[cfg(feature = "serde")]
4use serde::{Deserialize, Serialize};
5
6use crate::SignalDirection as Direction;
7
8#[derive(Debug, Clone, PartialEq)]
11pub struct DirectionalStatement<T> {
12 pub tf: T,
13 pub direction: Direction,
14 pub strength: f64,
15}
16
17#[derive(Debug, Clone, PartialEq)]
18pub struct Agreement {
19 pub direction: Direction,
20 pub agreement: f64,
30 pub conflict: bool,
35}
36
37#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
39#[cfg_attr(
40 feature = "serde",
41 derive(Serialize, Deserialize),
42 serde(rename_all = "snake_case")
43)]
44pub enum AgreementStrategy {
45 #[default]
48 Majority,
49 WeightedByStrength,
52 TimeframeTopDown,
57}
58
59pub fn aggregate_agreement<T: Copy + Eq>(
66 mut statements: Vec<DirectionalStatement<T>>,
67 strategies: &[AgreementStrategy],
68 tf_order: &[T],
69) -> Agreement {
70 let mut tally_by_strength = false;
71
72 for strategy in strategies {
73 match strategy {
74 AgreementStrategy::TimeframeTopDown => {
75 statements = apply_top_down(statements, tf_order);
76 }
77 AgreementStrategy::Majority => tally_by_strength = false,
78 AgreementStrategy::WeightedByStrength => tally_by_strength = true,
79 }
80 }
81
82 let result = tally(&statements, tally_by_strength);
83 if result.direction != Direction::Neutral {
84 return result;
85 }
86
87 let tie_broken = tally(
93 &apply_top_down(statements.clone(), tf_order),
94 tally_by_strength,
95 );
96 if tie_broken.direction != Direction::Neutral {
97 return tie_broken;
98 }
99
100 Agreement {
101 conflict: has_opposing_statements(&statements),
102 ..result
103 }
104}
105
106fn has_opposing_statements<T>(statements: &[DirectionalStatement<T>]) -> bool {
107 let (mut bull, mut bear) = (false, false);
108 for s in statements {
109 match s.direction {
110 Direction::Bullish => bull = true,
111 Direction::Bearish => bear = true,
112 Direction::Neutral => {}
113 }
114 }
115 bull && bear
116}
117
118fn apply_top_down<T: Copy + Eq>(
124 statements: Vec<DirectionalStatement<T>>,
125 tf_order: &[T],
126) -> Vec<DirectionalStatement<T>> {
127 let bias = tf_order.iter().find_map(|&tf| {
128 let on_tf: Vec<&DirectionalStatement<T>> =
129 statements.iter().filter(|s| s.tf == tf).collect();
130 if on_tf.is_empty() {
131 return None;
132 }
133 match tally_direction(on_tf.into_iter().cloned(), false) {
134 Direction::Neutral => None,
135 dir => Some(dir),
136 }
137 });
138
139 match bias {
140 None => statements,
141 Some(bias) => statements
142 .into_iter()
143 .filter(|s| s.direction == bias)
144 .collect(),
145 }
146}
147
148fn tally<T: Clone>(statements: &[DirectionalStatement<T>], by_strength: bool) -> Agreement {
149 let direction = tally_direction(statements.iter().cloned(), by_strength);
150
151 let (bull, bear): (f64, f64) = statements.iter().fold((0.0, 0.0), |(b, r), s| {
152 let weight = if by_strength { s.strength } else { 1.0 };
153 match s.direction {
154 Direction::Bullish => (b + weight, r),
155 Direction::Bearish => (b, r + weight),
156 Direction::Neutral => (b, r),
157 }
158 });
159 let total = bull + bear;
160 let agreement = if total > 0.0 {
161 (bull.max(bear)) / total
162 } else {
163 0.0
164 };
165
166 Agreement {
167 direction,
168 agreement,
169 conflict: false,
170 }
171}
172
173fn tally_direction<T>(
174 statements: impl Iterator<Item = DirectionalStatement<T>>,
175 by_strength: bool,
176) -> Direction {
177 let mut weight: HashMap<Direction, f64> = HashMap::new();
178 for s in statements {
179 if s.direction == Direction::Neutral {
180 continue;
181 }
182 let w = if by_strength { s.strength } else { 1.0 };
183 *weight.entry(s.direction).or_insert(0.0) += w;
184 }
185 let bull = weight.get(&Direction::Bullish).copied().unwrap_or(0.0);
186 let bear = weight.get(&Direction::Bearish).copied().unwrap_or(0.0);
187 if bull == 0.0 && bear == 0.0 {
188 Direction::Neutral
189 } else if bull > bear {
190 Direction::Bullish
191 } else if bear > bull {
192 Direction::Bearish
193 } else {
194 Direction::Neutral }
196}