1use crate::aes::Aesthetic;
2use crate::data::{DataFrame, Value};
3use crate::scale::ScaleSet;
4
5use super::Stat;
6
7pub struct StatBinHex {
10 pub bins_x: usize,
11 pub bins_y: usize,
12}
13
14impl Default for StatBinHex {
15 fn default() -> Self {
16 StatBinHex {
17 bins_x: 30,
18 bins_y: 30,
19 }
20 }
21}
22
23impl Stat for StatBinHex {
24 fn compute_group(&self, data: &DataFrame, _scales: &ScaleSet) -> DataFrame {
25 let x_col = match data.column("x") {
26 Some(c) => c,
27 None => return DataFrame::new(),
28 };
29 let y_col = match data.column("y") {
30 Some(c) => c,
31 None => return DataFrame::new(),
32 };
33
34 let xs: Vec<f64> = x_col.iter().filter_map(|v| v.as_f64()).collect();
35 let ys: Vec<f64> = y_col.iter().filter_map(|v| v.as_f64()).collect();
36 let n = xs.len().min(ys.len());
37 if n == 0 {
38 return DataFrame::new();
39 }
40
41 let x_min = xs.iter().cloned().fold(f64::INFINITY, f64::min);
42 let x_max = xs.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
43 let y_min = ys.iter().cloned().fold(f64::INFINITY, f64::min);
44 let y_max = ys.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
45
46 let x_range = if (x_max - x_min).abs() < f64::EPSILON {
47 1.0
48 } else {
49 x_max - x_min
50 };
51 let y_range = if (y_max - y_min).abs() < f64::EPSILON {
52 1.0
53 } else {
54 y_max - y_min
55 };
56
57 let hex_w = x_range / self.bins_x as f64;
59 let hex_h = y_range / self.bins_y as f64;
60
61 let mut counts: std::collections::BTreeMap<(i64, i64), usize> =
65 std::collections::BTreeMap::new();
66
67 for i in 0..n {
68 let col = ((xs[i] - x_min) / hex_w).floor() as i64;
70 let row = ((ys[i] - y_min) / hex_h).floor() as i64;
71
72 let adj_col = if row % 2 != 0 {
74 ((xs[i] - x_min - hex_w * 0.5) / hex_w).floor() as i64
75 } else {
76 col
77 };
78
79 *counts.entry((row, adj_col)).or_insert(0) += 1;
80 }
81
82 let mut x_vals = Vec::new();
83 let mut y_vals = Vec::new();
84 let mut fill_vals = Vec::new();
85
86 for (&(row, col), &count) in &counts {
87 if count == 0 {
88 continue;
89 }
90 let cx =
92 x_min + (col as f64 + 0.5) * hex_w + if row % 2 != 0 { hex_w * 0.5 } else { 0.0 };
93 let cy = y_min + (row as f64 + 0.5) * hex_h;
94
95 x_vals.push(Value::Float(cx));
96 y_vals.push(Value::Float(cy));
97 fill_vals.push(Value::Float(count as f64));
98 }
99
100 let mut result = DataFrame::new();
101 result.add_column("x".to_string(), x_vals);
102 result.add_column("y".to_string(), y_vals);
103 result.add_column("fill".to_string(), fill_vals);
104
105 result
106 }
107
108 fn required_aes(&self) -> Vec<Aesthetic> {
109 vec![Aesthetic::X, Aesthetic::Y]
110 }
111
112 fn name(&self) -> &str {
113 "binhex"
114 }
115}
116
117#[cfg(test)]
118mod tests {
119 use super::*;
120
121 #[test]
122 fn test_binhex_basic() {
123 let mut data = DataFrame::new();
124 let x_vals: Vec<Value> = (0..100).map(|i| Value::Float(i as f64 / 10.0)).collect();
125 let y_vals: Vec<Value> = (0..100).map(|i| Value::Float(i as f64 / 5.0)).collect();
126 data.add_column("x".to_string(), x_vals);
127 data.add_column("y".to_string(), y_vals);
128
129 let stat = StatBinHex {
130 bins_x: 5,
131 bins_y: 5,
132 };
133 let scales = ScaleSet::new();
134 let result = stat.compute_group(&data, &scales);
135
136 assert!(result.nrows() > 0);
137 assert!(result.column("x").is_some());
138 assert!(result.column("y").is_some());
139 assert!(result.column("fill").is_some());
140 }
141
142 #[test]
143 fn output_order_is_deterministic() {
144 let mut data = DataFrame::new();
145 let x: Vec<Value> = (0..500)
146 .map(|i| Value::Float(((i * 37) % 101) as f64))
147 .collect();
148 let y: Vec<Value> = (0..500)
149 .map(|i| Value::Float(((i * 53) % 97) as f64))
150 .collect();
151 data.add_column("x".to_string(), x);
152 data.add_column("y".to_string(), y);
153 let stat = StatBinHex {
154 bins_x: 12,
155 bins_y: 12,
156 };
157 let scales = ScaleSet::new();
158 let a = stat.compute_group(&data, &scales);
159 for _ in 0..5 {
160 let b = stat.compute_group(&data, &scales);
161 assert_eq!(a.column("x"), b.column("x"));
162 assert_eq!(a.column("y"), b.column("y"));
163 assert_eq!(a.column("fill"), b.column("fill"));
164 }
165 let ys: Vec<f64> = a
167 .column("y")
168 .unwrap()
169 .iter()
170 .map(|v| v.as_f64().unwrap())
171 .collect();
172 assert!(ys.windows(2).all(|w| w[0] <= w[1]));
173 let total: f64 = a
174 .column("fill")
175 .unwrap()
176 .iter()
177 .map(|v| v.as_f64().unwrap())
178 .sum();
179 assert_eq!(total, 500.0);
180 }
181}