1use std::collections::VecDeque;
4
5use crate::error::{Error, Result};
6use crate::traits::Indicator;
7
8fn sign(a: f64, b: f64) -> i32 {
10 if a > b {
11 1
12 } else if a < b {
13 -1
14 } else {
15 0
16 }
17}
18
19#[derive(Debug, Clone)]
58pub struct KendallTau {
59 period: usize,
60 window: VecDeque<(f64, f64)>,
61 counts: PairCounts,
66 last: Option<f64>,
67}
68
69#[derive(Debug, Clone, Copy, Default)]
71struct PairCounts {
72 concordant: i64,
73 discordant: i64,
74 tie_x: i64,
75 tie_y: i64,
76}
77
78impl PairCounts {
79 fn apply(&mut self, a: (f64, f64), b: (f64, f64), step: i64) {
81 let sx = sign(b.0, a.0);
82 let sy = sign(b.1, a.1);
83 if sx == 0 {
84 self.tie_x += step;
85 }
86 if sy == 0 {
87 self.tie_y += step;
88 }
89 let prod = sx * sy;
90 if prod > 0 {
91 self.concordant += step;
92 } else if prod < 0 {
93 self.discordant += step;
94 }
95 }
96}
97
98impl KendallTau {
99 pub fn new(period: usize) -> Result<Self> {
106 if period < 2 {
107 return Err(Error::InvalidPeriod {
108 message: "Kendall tau needs period >= 2",
109 });
110 }
111 if period > crate::error::MAX_PERIOD {
112 return Err(Error::InvalidPeriod {
113 message: crate::error::PERIOD_ABOVE_MAX,
114 });
115 }
116 Ok(Self {
117 period,
118 window: VecDeque::with_capacity(period),
119 counts: PairCounts::default(),
120 last: None,
121 })
122 }
123
124 pub const fn period(&self) -> usize {
126 self.period
127 }
128
129 pub const fn value(&self) -> Option<f64> {
131 self.last
132 }
133
134 fn compute(&self) -> f64 {
135 let len = self.window.len();
136 let PairCounts {
137 concordant,
138 discordant,
139 tie_x,
140 tie_y,
141 } = self.counts;
142 let n0 = (len * (len - 1) / 2) as f64;
143 let denom = ((n0 - tie_x as f64) * (n0 - tie_y as f64)).sqrt();
144 if denom == 0.0 {
145 return 0.0;
146 }
147 ((concordant - discordant) as f64 / denom).clamp(-1.0, 1.0)
148 }
149}
150
151impl Indicator for KendallTau {
152 type Input = (f64, f64);
153 type Output = f64;
154
155 #[inline]
156 fn update(&mut self, input: (f64, f64)) -> Option<f64> {
157 if !input.0.is_finite() || !input.1.is_finite() {
158 return None;
159 }
160 if self.window.len() == self.period {
161 let oldest = self.window.pop_front().expect("window is full");
162 for &other in &self.window {
163 self.counts.apply(oldest, other, -1);
164 }
165 }
166 for &other in &self.window {
167 self.counts.apply(other, input, 1);
168 }
169 self.window.push_back(input);
170 if self.window.len() < self.period {
171 return None;
172 }
173 let out = self.compute();
174 self.last = Some(out);
175 Some(out)
176 }
177
178 fn reset(&mut self) {
179 self.window.clear();
180 self.counts = PairCounts::default();
181 self.last = None;
182 }
183
184 #[inline]
185 fn warmup_period(&self) -> usize {
186 self.period
187 }
188
189 #[inline]
190 fn is_ready(&self) -> bool {
191 self.last.is_some()
192 }
193
194 #[inline]
195 fn name(&self) -> &'static str {
196 "KendallTau"
197 }
198}
199
200#[cfg(test)]
201mod tests {
202 use super::*;
203 use crate::traits::BatchExt;
204 use approx::assert_relative_eq;
205
206 #[test]
207 fn rejects_period_below_two() {
208 assert!(matches!(
209 KendallTau::new(1),
210 Err(Error::InvalidPeriod { .. })
211 ));
212 assert!(KendallTau::new(2).is_ok());
213 }
214
215 #[test]
216 fn accessors_and_metadata() {
217 let k = KendallTau::new(20).unwrap();
218 assert_eq!(k.period(), 20);
219 assert_eq!(k.warmup_period(), 20);
220 assert_eq!(k.name(), "KendallTau");
221 assert!(!k.is_ready());
222 assert_eq!(k.value(), None);
223 }
224
225 #[test]
226 fn first_emission_at_warmup_period() {
227 let mut k = KendallTau::new(4).unwrap();
228 let out = k.batch(&[(1.0, 1.0), (2.0, 2.0), (3.0, 3.0), (4.0, 4.0), (5.0, 5.0)]);
229 for v in out.iter().take(3) {
230 assert!(v.is_none());
231 }
232 assert!(out[3].is_some());
233 }
234
235 #[test]
236 fn monotone_increasing_is_one() {
237 let pairs: Vec<(f64, f64)> = (0..20)
238 .map(|i| (f64::from(i), 2.0 * f64::from(i) + 1.0))
239 .collect();
240 let last = KendallTau::new(10)
241 .unwrap()
242 .batch(&pairs)
243 .into_iter()
244 .flatten()
245 .last()
246 .unwrap();
247 assert_relative_eq!(last, 1.0, epsilon = 1e-9);
248 }
249
250 #[test]
251 fn monotone_decreasing_is_minus_one() {
252 let pairs: Vec<(f64, f64)> = (0..20)
253 .map(|i| (f64::from(i), -3.0 * f64::from(i)))
254 .collect();
255 let last = KendallTau::new(10)
256 .unwrap()
257 .batch(&pairs)
258 .into_iter()
259 .flatten()
260 .last()
261 .unwrap();
262 assert_relative_eq!(last, -1.0, epsilon = 1e-9);
263 }
264
265 #[test]
266 fn constant_channel_yields_zero() {
267 let pairs: Vec<(f64, f64)> = (0..20).map(|i| (f64::from(i), 7.0)).collect();
269 let last = KendallTau::new(8)
270 .unwrap()
271 .batch(&pairs)
272 .into_iter()
273 .flatten()
274 .last()
275 .unwrap();
276 assert_relative_eq!(last, 0.0, epsilon = 1e-12);
277 }
278
279 #[test]
280 fn output_in_range() {
281 let pairs: Vec<(f64, f64)> = (0..80)
282 .map(|i| {
283 let t = f64::from(i);
284 (100.0 + t.sin() * 5.0, 50.0 + (t * 0.3).cos() * 3.0)
285 })
286 .collect();
287 for v in KendallTau::new(20)
288 .unwrap()
289 .batch(&pairs)
290 .into_iter()
291 .flatten()
292 {
293 assert!((-1.0..=1.0).contains(&v));
294 }
295 }
296
297 #[test]
298 fn reset_clears_state() {
299 let mut k = KendallTau::new(4).unwrap();
300 k.batch(&[(1.0, 1.0), (2.0, 2.0), (3.0, 3.0), (4.0, 4.0)]);
301 assert!(k.is_ready());
302 k.reset();
303 assert!(!k.is_ready());
304 assert_eq!(k.value(), None);
305 assert_eq!(k.update((1.0, 1.0)), None);
306 }
307
308 #[test]
309 fn batch_equals_streaming() {
310 let pairs: Vec<(f64, f64)> = (0..60)
311 .map(|i| {
312 let t = f64::from(i);
313 (t.sin(), (t * 0.5).cos())
314 })
315 .collect();
316 let batch = KendallTau::new(14).unwrap().batch(&pairs);
317 let mut b = KendallTau::new(14).unwrap();
318 let streamed: Vec<_> = pairs.iter().map(|p| b.update(*p)).collect();
319 assert_eq!(batch, streamed);
320 }
321
322 #[test]
323 fn ties_are_corrected() {
324 let mut k = KendallTau::new(4).unwrap();
327 assert_eq!(k.update((1.0, 1.0)), None);
328 assert_eq!(k.update((1.0, 2.0)), None);
329 assert_eq!(k.update((2.0, 2.0)), None);
330 let v = k.update((3.0, 3.0)).unwrap();
331 assert!((-1.0..=1.0).contains(&v), "got {v}");
332 }
333
334 #[test]
335 fn non_finite_input_returns_none() {
336 let mut k = KendallTau::new(2).unwrap();
337 assert_eq!(k.update((f64::NAN, 1.0)), None);
338 assert_eq!(k.update((1.0, f64::INFINITY)), None);
339 assert_eq!(k.update((1.0, 2.0)), None);
341 assert!(k.update((2.0, 5.0)).is_some());
342 }
343}