1use std::collections::HashMap;
11use std::collections::VecDeque;
12
13pub fn pearson_correlation(x: &[f64], y: &[f64]) -> Option<f64> {
22 if x.len() != y.len() || x.len() < 2 {
23 return None;
24 }
25 let n = x.len() as f64;
26 let sum_x: f64 = x.iter().sum();
27 let sum_y: f64 = y.iter().sum();
28 let sum_xy: f64 = x.iter().zip(y.iter()).map(|(a, b)| a * b).sum();
29 let sum_x2: f64 = x.iter().map(|a| a * a).sum();
30 let sum_y2: f64 = y.iter().map(|b| b * b).sum();
31
32 let num = n * sum_xy - sum_x * sum_y;
33 let den_sq = (n * sum_x2 - sum_x * sum_x) * (n * sum_y2 - sum_y * sum_y);
34 if den_sq <= 0.0 {
35 return None;
36 }
37 Some((num / den_sq.sqrt()).clamp(-1.0, 1.0))
38}
39
40pub fn spearman_correlation(x: &[f64], y: &[f64]) -> Option<f64> {
47 if x.len() != y.len() || x.len() < 2 {
48 return None;
49 }
50 let rx = average_ranks(x);
51 let ry = average_ranks(y);
52 pearson_correlation(&rx, &ry)
53}
54
55fn average_ranks(data: &[f64]) -> Vec<f64> {
57 let n = data.len();
58 let mut indexed: Vec<(f64, usize)> = data.iter().copied().enumerate().map(|(i, v)| (v, i)).collect();
60 indexed.sort_by(|a, b| a.0.partial_cmp(&b.0).unwrap_or(std::cmp::Ordering::Equal));
61
62 let mut ranks = vec![0.0_f64; n];
63 let mut i = 0;
64 while i < n {
65 let mut j = i + 1;
67 while j < n && (indexed[j].0 - indexed[i].0).abs() < f64::EPSILON {
68 j += 1;
69 }
70 let avg_rank = (i + j + 1) as f64 / 2.0; for k in i..j {
73 ranks[indexed[k].1] = avg_rank;
74 }
75 i = j;
76 }
77 ranks
78}
79
80pub fn kendall_tau(x: &[f64], y: &[f64]) -> Option<f64> {
87 if x.len() != y.len() || x.len() < 2 {
88 return None;
89 }
90 let n = x.len();
91 let mut concordant: i64 = 0;
92 let mut discordant: i64 = 0;
93 let mut ties_x: i64 = 0;
94 let mut ties_y: i64 = 0;
95 let mut ties_xy: i64 = 0;
96
97 for i in 0..n {
98 for j in (i + 1)..n {
99 let dx = x[i] - x[j];
100 let dy = y[i] - y[j];
101 let prod = dx * dy;
102 let x_tied = dx.abs() < f64::EPSILON;
103 let y_tied = dy.abs() < f64::EPSILON;
104
105 if x_tied && y_tied {
106 ties_xy += 1;
107 } else if x_tied {
108 ties_x += 1;
109 } else if y_tied {
110 ties_y += 1;
111 } else if prod > 0.0 {
112 concordant += 1;
113 } else {
114 discordant += 1;
115 }
116 }
117 }
118
119 let total_pairs = (n as i64 * (n as i64 - 1)) / 2;
120 let n0 = total_pairs;
121 let n1 = n0 - ties_x - ties_xy;
122 let n2 = n0 - ties_y - ties_xy;
123
124 let denom = (n1 as f64 * n2 as f64).sqrt();
125 if denom == 0.0 {
126 return None;
127 }
128
129 let tau = (concordant - discordant) as f64 / denom;
130 Some(tau.clamp(-1.0, 1.0))
131}
132
133#[derive(Debug, Clone)]
140pub struct SymbolCorrelationMatrix {
141 pub symbols: Vec<String>,
143 pub matrix: Vec<Vec<f64>>,
145 pub n: usize,
147}
148
149impl SymbolCorrelationMatrix {
150 pub fn from_returns(symbols: Vec<String>, returns: Vec<Vec<f64>>) -> Self {
155 let n = symbols.len();
156 let mut matrix = vec![vec![1.0_f64; n]; n];
157
158 for i in 0..n {
159 for j in (i + 1)..n {
160 let corr = pearson_correlation(&returns[i], &returns[j]).unwrap_or(0.0);
161 matrix[i][j] = corr;
162 matrix[j][i] = corr;
163 }
164 }
165
166 Self { symbols, matrix, n }
167 }
168
169 pub fn get(&self, i: usize, j: usize) -> f64 {
171 self.matrix[i][j]
172 }
173
174 pub fn to_table(&self) -> String {
176 let col_w = self.symbols.iter().map(|s| s.len()).max().unwrap_or(6).max(6);
178 let fmt = |v: f64| format!("{:>width$.4}", v, width = col_w);
179 let pad = |s: &str| format!("{:>width$}", s, width = col_w);
180
181 let mut out = String::new();
182 out.push_str(&" ".repeat(col_w + 1));
184 for sym in &self.symbols {
185 out.push(' ');
186 out.push_str(&pad(sym));
187 }
188 out.push('\n');
189
190 for (i, sym) in self.symbols.iter().enumerate() {
191 out.push_str(&pad(sym));
192 for j in 0..self.n {
193 out.push(' ');
194 out.push_str(&fmt(self.matrix[i][j]));
195 }
196 out.push('\n');
197 }
198 out
199 }
200
201 pub fn highly_correlated(&self, threshold: f64) -> Vec<(String, String, f64)> {
205 let mut result = Vec::new();
206 for i in 0..self.n {
207 for j in (i + 1)..self.n {
208 let c = self.matrix[i][j];
209 if c.abs() > threshold {
210 result.push((self.symbols[i].clone(), self.symbols[j].clone(), c));
211 }
212 }
213 }
214 result
215 }
216
217 pub fn eigenvalues(&self) -> Vec<f64> {
222 if self.n == 0 {
223 return vec![];
224 }
225 jacobi_eigenvalues(&self.matrix, self.n)
226 }
227}
228
229fn jacobi_eigenvalues(matrix: &[Vec<f64>], n: usize) -> Vec<f64> {
234 let mut a: Vec<f64> = matrix.iter().flat_map(|row| row.iter().copied()).collect();
236 let idx = |i: usize, j: usize| i * n + j;
237
238 let max_sweeps = 100;
239 let tol = 1e-10_f64;
240
241 for _ in 0..max_sweeps {
242 let mut max_val = 0.0_f64;
244 for i in 0..n {
245 for j in (i + 1)..n {
246 let v = a[idx(i, j)].abs();
247 if v > max_val {
248 max_val = v;
249 }
250 }
251 }
252 if max_val < tol {
253 break;
254 }
255
256 for p in 0..n {
258 for q in (p + 1)..n {
259 let apq = a[idx(p, q)];
260 if apq.abs() < tol {
261 continue;
262 }
263 let app = a[idx(p, p)];
264 let aqq = a[idx(q, q)];
265 let theta = 0.5 * (aqq - app) / apq;
266 let t = if theta >= 0.0 {
267 1.0 / (theta + (1.0 + theta * theta).sqrt())
268 } else {
269 -1.0 / (-theta + (1.0 + theta * theta).sqrt())
270 };
271 let c = 1.0 / (1.0 + t * t).sqrt();
272 let s = t * c;
273
274 a[idx(p, p)] = app - t * apq;
276 a[idx(q, q)] = aqq + t * apq;
277 a[idx(p, q)] = 0.0;
278 a[idx(q, p)] = 0.0;
279
280 for r in 0..n {
282 if r == p || r == q {
283 continue;
284 }
285 let arp = a[idx(r, p)];
286 let arq = a[idx(r, q)];
287 a[idx(r, p)] = c * arp - s * arq;
288 a[idx(p, r)] = a[idx(r, p)];
289 a[idx(r, q)] = s * arp + c * arq;
290 a[idx(q, r)] = a[idx(r, q)];
291 }
292 }
293 }
294 }
295
296 let mut eigs: Vec<f64> = (0..n).map(|i| a[idx(i, i)]).collect();
298 eigs.sort_by(|a, b| b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal));
299 eigs
300}
301
302pub struct RollingCorrelation {
309 window: usize,
311 series: HashMap<String, VecDeque<f64>>,
313}
314
315impl RollingCorrelation {
316 pub fn new(window: usize) -> Self {
318 Self { window, series: HashMap::new() }
319 }
320
321 pub fn push(&mut self, symbol: &str, value: f64) {
326 let dq = self.series.entry(symbol.to_string()).or_insert_with(|| VecDeque::with_capacity(self.window));
327 if dq.len() >= self.window {
328 dq.pop_front();
329 }
330 dq.push_back(value);
331 }
332
333 pub fn is_ready(&self) -> bool {
335 !self.series.is_empty() && self.series.values().all(|dq| dq.len() >= self.window)
336 }
337
338 pub fn compute_matrix(&self) -> Option<SymbolCorrelationMatrix> {
342 if !self.is_ready() {
343 return None;
344 }
345 let mut symbols: Vec<String> = self.series.keys().cloned().collect();
346 symbols.sort();
347 let returns: Vec<Vec<f64>> = symbols
348 .iter()
349 .map(|s| self.series[s].iter().copied().collect())
350 .collect();
351 Some(SymbolCorrelationMatrix::from_returns(symbols, returns))
352 }
353
354 pub fn pairwise(&self, sym_a: &str, sym_b: &str) -> Option<f64> {
358 let a = self.series.get(sym_a)?;
359 let b = self.series.get(sym_b)?;
360 if a.len() < self.window || b.len() < self.window {
361 return None;
362 }
363 let va: Vec<f64> = a.iter().copied().collect();
364 let vb: Vec<f64> = b.iter().copied().collect();
365 pearson_correlation(&va, &vb)
366 }
367}
368
369#[cfg(test)]
372mod tests {
373 use super::*;
374
375 #[test]
376 fn test_pearson_perfect_positive() {
377 let x = vec![1.0, 2.0, 3.0, 4.0, 5.0];
378 let y = vec![2.0, 4.0, 6.0, 8.0, 10.0];
379 let r = pearson_correlation(&x, &y).unwrap();
380 assert!((r - 1.0).abs() < 1e-9, "r={r}");
381 }
382
383 #[test]
384 fn test_pearson_perfect_negative() {
385 let x = vec![1.0, 2.0, 3.0, 4.0, 5.0];
386 let y = vec![10.0, 8.0, 6.0, 4.0, 2.0];
387 let r = pearson_correlation(&x, &y).unwrap();
388 assert!((r + 1.0).abs() < 1e-9, "r={r}");
389 }
390
391 #[test]
392 fn test_pearson_insufficient_data() {
393 assert!(pearson_correlation(&[1.0], &[1.0]).is_none());
394 assert!(pearson_correlation(&[], &[]).is_none());
395 }
396
397 #[test]
398 fn test_pearson_zero_variance() {
399 let x = vec![5.0, 5.0, 5.0];
400 let y = vec![1.0, 2.0, 3.0];
401 assert!(pearson_correlation(&x, &y).is_none());
402 }
403
404 #[test]
405 fn test_spearman_perfect_positive() {
406 let x = vec![1.0, 2.0, 3.0, 4.0, 5.0];
407 let y = vec![10.0, 20.0, 30.0, 40.0, 50.0];
408 let r = spearman_correlation(&x, &y).unwrap();
409 assert!((r - 1.0).abs() < 1e-9, "r={r}");
410 }
411
412 #[test]
413 fn test_spearman_anti_correlation() {
414 let x = vec![1.0, 2.0, 3.0, 4.0, 5.0];
415 let y = vec![5.0, 4.0, 3.0, 2.0, 1.0];
416 let r = spearman_correlation(&x, &y).unwrap();
417 assert!((r + 1.0).abs() < 1e-9, "r={r}");
418 }
419
420 #[test]
421 fn test_spearman_rank_transform_with_ties() {
422 let ranks = average_ranks(&[1.0, 1.0, 2.0]);
424 assert!((ranks[0] - 1.5).abs() < 1e-9);
425 assert!((ranks[1] - 1.5).abs() < 1e-9);
426 assert!((ranks[2] - 3.0).abs() < 1e-9);
427 }
428
429 #[test]
430 fn test_kendall_perfect_concordant() {
431 let x = vec![1.0, 2.0, 3.0, 4.0, 5.0];
432 let y = vec![1.0, 2.0, 3.0, 4.0, 5.0];
433 let tau = kendall_tau(&x, &y).unwrap();
434 assert!((tau - 1.0).abs() < 1e-9, "tau={tau}");
435 }
436
437 #[test]
438 fn test_kendall_perfect_discordant() {
439 let x = vec![1.0, 2.0, 3.0, 4.0, 5.0];
440 let y = vec![5.0, 4.0, 3.0, 2.0, 1.0];
441 let tau = kendall_tau(&x, &y).unwrap();
442 assert!((tau + 1.0).abs() < 1e-9, "tau={tau}");
443 }
444
445 #[test]
446 fn test_symbol_correlation_matrix_from_returns() {
447 let symbols = vec!["A".to_string(), "B".to_string(), "C".to_string()];
448 let returns = vec![
449 vec![1.0, 2.0, 3.0, 4.0, 5.0],
450 vec![2.0, 4.0, 6.0, 8.0, 10.0], vec![5.0, 4.0, 3.0, 2.0, 1.0], ];
453 let mat = SymbolCorrelationMatrix::from_returns(symbols, returns);
454 assert!((mat.get(0, 1) - 1.0).abs() < 1e-9);
455 assert!((mat.get(0, 2) + 1.0).abs() < 1e-9);
456 assert_eq!(mat.get(0, 0), 1.0);
457 }
458
459 #[test]
460 fn test_highly_correlated_filter() {
461 let symbols = vec!["A".to_string(), "B".to_string(), "C".to_string()];
462 let returns = vec![
463 vec![1.0, 2.0, 3.0, 4.0, 5.0],
464 vec![2.0, 4.0, 6.0, 8.0, 10.0],
465 vec![5.0, 4.0, 3.0, 2.0, 1.0],
466 ];
467 let mat = SymbolCorrelationMatrix::from_returns(symbols, returns);
468 let high = mat.highly_correlated(0.9);
469 assert_eq!(high.len(), 3);
471 }
472
473 #[test]
474 fn test_highly_correlated_excludes_below_threshold() {
475 let symbols = vec!["A".to_string(), "B".to_string()];
476 let returns = vec![
477 vec![1.0, 2.0, 3.0, 4.0, 5.0],
478 vec![1.0, 1.5, 1.0, 1.5, 1.0], ];
480 let mat = SymbolCorrelationMatrix::from_returns(symbols, returns);
481 let high = mat.highly_correlated(0.99);
482 assert!(high.is_empty());
483 }
484
485 #[test]
486 fn test_to_table_contains_symbols() {
487 let symbols = vec!["BTC".to_string(), "ETH".to_string()];
488 let returns = vec![
489 vec![1.0, 2.0, 3.0],
490 vec![1.0, 2.0, 3.0],
491 ];
492 let mat = SymbolCorrelationMatrix::from_returns(symbols, returns);
493 let table = mat.to_table();
494 assert!(table.contains("BTC"));
495 assert!(table.contains("ETH"));
496 }
497
498 #[test]
499 fn test_rolling_correlation_not_ready_until_window() {
500 let mut rc = RollingCorrelation::new(5);
501 for i in 0..4 {
502 rc.push("A", i as f64);
503 rc.push("B", i as f64 * 2.0);
504 }
505 assert!(!rc.is_ready());
506 assert!(rc.compute_matrix().is_none());
507 assert!(rc.pairwise("A", "B").is_none());
508 }
509
510 #[test]
511 fn test_rolling_correlation_ready_after_window() {
512 let mut rc = RollingCorrelation::new(5);
513 for i in 0..5 {
514 rc.push("A", i as f64);
515 rc.push("B", i as f64 * 2.0);
516 }
517 assert!(rc.is_ready());
518 let r = rc.pairwise("A", "B").unwrap();
519 assert!((r - 1.0).abs() < 1e-9, "r={r}");
520 }
521
522 #[test]
523 fn test_rolling_window_evicts_old_values() {
524 let mut rc = RollingCorrelation::new(3);
525 for i in 0..5 {
527 rc.push("A", i as f64);
528 }
529 let dq = &rc.series["A"];
530 assert_eq!(dq.len(), 3);
531 assert_eq!(dq[0], 2.0);
532 assert_eq!(dq[2], 4.0);
533 }
534
535 #[test]
536 fn test_rolling_correlation_matrix() {
537 let mut rc = RollingCorrelation::new(5);
538 for i in 0..5 {
539 let v = i as f64;
540 rc.push("X", v);
541 rc.push("Y", -v);
542 }
543 let mat = rc.compute_matrix().unwrap();
544 let x_idx = mat.symbols.iter().position(|s| s == "X").unwrap();
546 let y_idx = mat.symbols.iter().position(|s| s == "Y").unwrap();
547 assert!((mat.get(x_idx, y_idx) + 1.0).abs() < 1e-9);
548 }
549
550 #[test]
551 fn test_eigenvalues_length() {
552 let symbols = vec!["A".to_string(), "B".to_string(), "C".to_string()];
553 let returns = vec![
554 vec![1.0, 2.0, 3.0, 4.0, 5.0],
555 vec![2.0, 4.0, 6.0, 8.0, 10.0],
556 vec![5.0, 4.0, 3.0, 2.0, 1.0],
557 ];
558 let mat = SymbolCorrelationMatrix::from_returns(symbols, returns);
559 let eigs = mat.eigenvalues();
560 assert_eq!(eigs.len(), 3);
561 }
562
563 #[test]
564 fn test_eigenvalues_descending() {
565 let symbols = vec!["A".to_string(), "B".to_string(), "C".to_string()];
566 let returns = vec![
567 vec![1.0, 2.0, 3.0, 4.0, 5.0],
568 vec![5.0, 3.0, 1.0, 4.0, 2.0],
569 vec![2.0, 5.0, 1.0, 3.0, 4.0],
570 ];
571 let mat = SymbolCorrelationMatrix::from_returns(symbols, returns);
572 let eigs = mat.eigenvalues();
573 for i in 0..eigs.len() - 1 {
574 assert!(eigs[i] >= eigs[i + 1] - 1e-9, "eigs not descending: {:?}", eigs);
575 }
576 }
577}