1use std::{collections::HashMap, hash::Hash};
4
5use num_traits::Zero;
6
7use super::Scale;
8
9#[derive(Clone)]
10pub struct ScaleBand<T> {
11 indices: HashMap<T, usize>,
17 band_count: usize,
19 max_band_width: Option<f32>,
21 range_start: f32,
22 range_diff: f32,
23 padding_inner: f32,
24 padding_outer: f32,
25}
26
27impl<T> ScaleBand<T> {
28 pub fn new(domain: impl IntoIterator<Item = T>, range: [f32; 2]) -> Self
31 where
32 T: Eq + Hash,
33 {
34 let mut indices = HashMap::new();
35 for value in domain {
36 let next = indices.len();
37 indices.entry(value).or_insert(next);
38 }
39
40 Self {
41 indices,
42 band_count: 0,
43 max_band_width: None,
44 range_start: range[0].min(range[1]),
45 range_diff: (range[1] - range[0]).abs(),
46 padding_inner: 0.,
47 padding_outer: 0.,
48 }
49 }
50
51 pub fn band_width(&self) -> f32 {
54 let width = self.avg_width() * (1. - self.padding_inner);
55 self.max_band_width
56 .map_or(width, |max_band_width| width.min(max_band_width))
57 }
58
59 pub fn max_band_width(mut self, width: f32) -> Self {
62 self.max_band_width = Some(width);
63 self
64 }
65
66 pub fn step(&self) -> f32 {
69 if self.len() <= 1 {
70 self.range_diff
71 } else {
72 self.display_avg_width() * self.ratio()
73 }
74 }
75
76 pub fn band_count(mut self, count: usize) -> Self {
80 self.band_count = count;
81 self
82 }
83
84 pub fn padding_inner(mut self, padding_inner: f32) -> Self {
86 self.padding_inner = padding_inner;
87 self
88 }
89
90 pub fn padding_outer(mut self, padding_outer: f32) -> Self {
92 self.padding_outer = padding_outer;
93 self
94 }
95
96 fn len(&self) -> usize {
99 self.indices.len().max(self.band_count)
100 }
101
102 fn avg_width(&self) -> f32 {
104 let len = self.len() as f32;
105 if len.is_zero() {
106 0.
107 } else {
108 self.range_diff / len
109 }
110 }
111
112 fn ratio(&self) -> f32 {
114 1. + self.padding_inner / (self.len() - 1) as f32
115 }
116
117 fn display_avg_width(&self) -> f32 {
119 let padding_outer_width = self.avg_width() * self.padding_outer;
120 (self.range_diff - padding_outer_width * 2.) / self.len() as f32
121 }
122}
123
124impl<T> Scale<T> for ScaleBand<T>
125where
126 T: Eq + Hash,
127{
128 fn tick(&self, value: &T) -> Option<f32> {
129 let index = *self.indices.get(value)?;
130 let domain_len = self.len();
131
132 if domain_len == 1 {
134 return Some(self.range_start + (self.range_diff - self.band_width()) / 2.);
135 }
136
137 let avg_width = self.display_avg_width();
138 let padding_outer_width = self.avg_width() * self.padding_outer;
139 Some(self.range_start + index as f32 * avg_width * self.ratio() + padding_outer_width)
140 }
141
142 fn nearest_index(&self, tick: f32) -> usize {
143 let domain_len = self.len();
144 if domain_len == 0 {
145 return 0;
146 }
147
148 if domain_len == 1 {
150 return 0;
151 }
152
153 let avg_width = self.display_avg_width();
154 let padding_outer_width = self.avg_width() * self.padding_outer;
155 let adjusted_tick = tick - self.range_start - padding_outer_width;
156 let index = (adjusted_tick / (avg_width * self.ratio())).round() as i32;
157
158 (index.max(0) as usize).min(domain_len.saturating_sub(1))
159 }
160}
161
162#[cfg(test)]
163mod tests {
164 use super::*;
165
166 #[test]
167 fn test_scale_band() {
168 let scale = ScaleBand::new(vec![1, 2, 3], [0., 90.]);
169 assert_eq!(scale.tick(&1), Some(0.));
170 assert_eq!(scale.tick(&2), Some(30.));
171 assert_eq!(scale.tick(&3), Some(60.));
172 assert_eq!(scale.band_width(), 30.);
173 }
174
175 #[test]
176 fn max_band_width_caps_the_width_but_not_the_ticks() {
177 let wide = ScaleBand::new(vec![1, 2], [0., 200.]);
178 let capped = ScaleBand::new(vec![1, 2], [0., 200.]).max_band_width(30.);
179 assert_eq!(wide.band_width(), 100.);
180 assert_eq!(capped.band_width(), 30.);
181 assert_eq!(capped.tick(&2), wide.tick(&2));
182 }
183
184 #[test]
185 fn test_scale_band_dedup() {
186 let scale = ScaleBand::new(vec![1, 2, 3, 1, 2, 3], [0., 90.]);
188 assert_eq!(scale.len(), 3);
189 assert_eq!(scale.tick(&1), Some(0.));
190 assert_eq!(scale.tick(&2), Some(30.));
191 assert_eq!(scale.tick(&3), Some(60.));
192 assert_eq!(scale.band_width(), 30.);
193 }
194
195 #[test]
196 fn test_scale_band_step() {
197 let scale = ScaleBand::new(vec![1, 2, 3], [0., 90.]);
199 assert_eq!(
200 scale.step(),
201 scale.tick(&2).unwrap() - scale.tick(&1).unwrap()
202 );
203
204 let padded = ScaleBand::new(vec![1, 2, 3], [0., 90.])
205 .padding_inner(0.4)
206 .padding_outer(0.2);
207 assert!(
208 (padded.step() - (padded.tick(&2).unwrap() - padded.tick(&1).unwrap())).abs() < 1e-4
209 );
210
211 assert_eq!(ScaleBand::new(vec![1], [0., 90.]).step(), 90.);
213 }
214
215 #[test]
216 fn test_scale_band_count() {
217 let scale = |domain: Vec<i32>| {
218 ScaleBand::new(domain, [0., 100.])
219 .band_count(4)
220 .padding_inner(0.4)
221 .padding_outer(0.2)
222 };
223
224 let short = scale(vec![1, 2]);
226 let full = scale(vec![1, 2, 3, 4]);
227 assert_eq!(short.tick(&2), full.tick(&2));
228 assert_eq!(short.band_width(), full.band_width());
229 assert_eq!(short.step(), full.step());
230
231 assert_eq!(short.nearest_index(full.tick(&4).unwrap()), 3);
233
234 assert_eq!(scale(vec![1]).tick(&1), full.tick(&1));
236
237 let domain = ScaleBand::new(vec![1, 2, 3], [0., 90.]);
239 assert_eq!(domain.band_count(2).tick(&3), Some(60.));
240 }
241
242 #[test]
243 fn test_scale_band_zero() {
244 let scale = ScaleBand::new(vec![], [0., 90.]);
245 assert_eq!(scale.tick(&1), None);
246 assert_eq!(scale.tick(&2), None);
247 assert_eq!(scale.tick(&3), None);
248 assert_eq!(scale.band_width(), 0.);
249
250 let scale = ScaleBand::new(vec![1, 2, 3], [0., 0.]);
251 assert_eq!(scale.tick(&1), Some(0.));
252 assert_eq!(scale.tick(&2), Some(0.));
253 assert_eq!(scale.tick(&3), Some(0.));
254 assert_eq!(scale.band_width(), 0.);
255 }
256
257 #[test]
258 fn test_scale_band_range_start() {
259 let scale = ScaleBand::new([1, 2, 3], [10., 100.]);
260 assert_eq!(scale.tick(&1), Some(10.));
261 assert_eq!(scale.tick(&2), Some(40.));
262 assert_eq!(scale.nearest_index(41.), 1);
263 assert_eq!(ScaleBand::new([1, 2, 3], [100., 10.]).tick(&1), Some(10.));
265 }
266}