Skip to main content

gpui_base/plot/scale/
band.rs

1// @reference: https://d3js.org/d3-scale/band
2
3use std::{collections::HashMap, hash::Hash};
4
5use num_traits::Zero;
6
7use super::Scale;
8
9#[derive(Clone)]
10pub struct ScaleBand<T> {
11    /// Each distinct domain value paired with its band index.
12    ///
13    /// D3 keys its band domain through an `InternMap`, so a repeated value
14    /// keeps the index of its first occurrence and the band count follows the
15    /// distinct values, not the entry count.
16    indices: HashMap<T, usize>,
17    /// The bands laid out when more than the domain's; see [`Self::band_count`].
18    band_count: usize,
19    /// The widest a band may be; see [`Self::max_band_width`].
20    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    /// Lay the distinct values of `domain` out as bands across `range`, from
29    /// its lower end in domain order.
30    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    /// The width of a band: the range divided among the bands, less the inner
52    /// padding, and no wider than [`Self::max_band_width`] when set.
53    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    /// Cap the band width at `width`, so a few bands in a wide range stay
60    /// narrow; a band still starts where it would uncapped. Unset by default.
61    pub fn max_band_width(mut self, width: f32) -> Self {
62        self.max_band_width = Some(width);
63        self
64    }
65
66    /// The distance between the starts of two adjacent bands: the band width
67    /// plus the inner padding. The whole range for a single band.
68    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    /// Lay the range out for `count` bands, the domain taking the leading ones
77    /// in order and the rest staying empty. A `count` below the domain's length
78    /// has no effect.
79    pub fn band_count(mut self, count: usize) -> Self {
80        self.band_count = count;
81        self
82    }
83
84    /// Set the padding inner of the band.
85    pub fn padding_inner(mut self, padding_inner: f32) -> Self {
86        self.padding_inner = padding_inner;
87        self
88    }
89
90    /// Set the padding outer of the band.
91    pub fn padding_outer(mut self, padding_outer: f32) -> Self {
92        self.padding_outer = padding_outer;
93        self
94    }
95
96    /// The number of bands: one per distinct domain value, or the
97    /// [`Self::band_count`] when larger.
98    fn len(&self) -> usize {
99        self.indices.len().max(self.band_count)
100    }
101
102    /// The range divided evenly among the bands.
103    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    /// Get the ratio of the band.
113    fn ratio(&self) -> f32 {
114        1. + self.padding_inner / (self.len() - 1) as f32
115    }
116
117    /// Get the average width of the band for display.
118    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        // When there's only one element, place it in the center.
133        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        // Handle single element case
149        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        // Simulates grouped bar chart: 2 series × 3 categories = 6 entries, 3 unique.
187        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        // Adjacent bands start one step apart, whatever the padding.
198        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        // A single band spans the range.
212        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        // The domain takes the leading bands, each placed as if all were full.
225        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        // An empty band resolves past the domain rather than to its last value.
232        assert_eq!(short.nearest_index(full.tick(&4).unwrap()), 3);
233
234        // A single value sits in the first band instead of the center.
235        assert_eq!(scale(vec![1]).tick(&1), full.tick(&1));
236
237        // A count below the domain's length has no effect.
238        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        // The lower end leads whichever way the range is written.
264        assert_eq!(ScaleBand::new([1, 2, 3], [100., 10.]).tick(&1), Some(10.));
265    }
266}