Skip to main content

gpui_component/plot/scale/
band.rs

1// @reference: https://d3js.org/d3-scale/band
2
3use std::{collections::HashMap, hash::Hash};
4
5use itertools::Itertools;
6use num_traits::Zero;
7
8use super::Scale;
9
10#[derive(Clone)]
11pub struct ScaleBand<T> {
12    /// Each distinct domain value paired with its band index.
13    ///
14    /// D3 keys its band domain through an `InternMap`, so a repeated value
15    /// keeps the index of its first occurrence and the band count follows the
16    /// distinct values, not the entry count.
17    indices: HashMap<T, usize>,
18    range_diff: f32,
19    avg_width: f32,
20    padding_inner: f32,
21    padding_outer: f32,
22}
23
24impl<T> ScaleBand<T> {
25    pub fn new(domain: Vec<T>, range: Vec<f32>) -> Self
26    where
27        T: Eq + Hash,
28    {
29        let mut indices = HashMap::with_capacity(domain.len());
30        for value in domain {
31            let next = indices.len();
32            indices.entry(value).or_insert(next);
33        }
34
35        let len = indices.len() as f32;
36        let range_diff = range
37            .iter()
38            .minmax()
39            .into_option()
40            .map_or(0., |(min, max)| max - min);
41
42        Self {
43            indices,
44            range_diff,
45            avg_width: if len.is_zero() { 0. } else { range_diff / len },
46            padding_inner: 0.,
47            padding_outer: 0.,
48        }
49    }
50
51    /// Get the width of the band.
52    pub fn band_width(&self) -> f32 {
53        (self.avg_width * (1. - self.padding_inner)).min(30.)
54    }
55
56    /// The distance between the starts of two adjacent bands: the band width
57    /// plus the inner padding. The whole range for a single band.
58    pub fn step(&self) -> f32 {
59        if self.len() <= 1 {
60            self.range_diff
61        } else {
62            self.display_avg_width() * self.ratio()
63        }
64    }
65
66    /// Set the padding inner of the band.
67    pub fn padding_inner(mut self, padding_inner: f32) -> Self {
68        self.padding_inner = padding_inner;
69        self
70    }
71
72    /// Set the padding outer of the band.
73    pub fn padding_outer(mut self, padding_outer: f32) -> Self {
74        self.padding_outer = padding_outer;
75        self
76    }
77
78    /// The number of bands, one per distinct domain value.
79    fn len(&self) -> usize {
80        self.indices.len()
81    }
82
83    /// Get the ratio of the band.
84    fn ratio(&self) -> f32 {
85        1. + self.padding_inner / (self.len() - 1) as f32
86    }
87
88    /// Get the average width of the band for display.
89    fn display_avg_width(&self) -> f32 {
90        let padding_outer_width = self.avg_width * self.padding_outer;
91        (self.range_diff - padding_outer_width * 2.) / self.len() as f32
92    }
93}
94
95impl<T> Scale<T> for ScaleBand<T>
96where
97    T: Eq + Hash,
98{
99    fn tick(&self, value: &T) -> Option<f32> {
100        let index = *self.indices.get(value)?;
101        let domain_len = self.len();
102
103        // When there's only one element, place it in the center.
104        if domain_len == 1 {
105            return Some((self.range_diff - self.band_width()) / 2.);
106        }
107
108        let avg_width = self.display_avg_width();
109        let padding_outer_width = self.avg_width * self.padding_outer;
110        Some(index as f32 * avg_width * self.ratio() + padding_outer_width)
111    }
112
113    fn least_index(&self, tick: f32) -> usize {
114        let domain_len = self.len();
115        if domain_len == 0 {
116            return 0;
117        }
118
119        // Handle single element case
120        if domain_len == 1 {
121            return 0;
122        }
123
124        let avg_width = self.display_avg_width();
125        let padding_outer_width = self.avg_width * self.padding_outer;
126        let adjusted_tick = tick - padding_outer_width;
127        let index = (adjusted_tick / (avg_width * self.ratio())).round() as i32;
128
129        (index.max(0) as usize).min(domain_len.saturating_sub(1))
130    }
131}
132
133#[cfg(test)]
134mod tests {
135    use super::*;
136
137    #[test]
138    fn test_scale_band() {
139        let scale = ScaleBand::new(vec![1, 2, 3], vec![0., 90.]);
140        assert_eq!(scale.tick(&1), Some(0.));
141        assert_eq!(scale.tick(&2), Some(30.));
142        assert_eq!(scale.tick(&3), Some(60.));
143        assert_eq!(scale.band_width(), 30.);
144    }
145
146    #[test]
147    fn test_scale_band_dedup() {
148        // Simulates grouped bar chart: 2 series × 3 categories = 6 entries, 3 unique.
149        let scale = ScaleBand::new(vec![1, 2, 3, 1, 2, 3], vec![0., 90.]);
150        assert_eq!(scale.len(), 3);
151        assert_eq!(scale.tick(&1), Some(0.));
152        assert_eq!(scale.tick(&2), Some(30.));
153        assert_eq!(scale.tick(&3), Some(60.));
154        assert_eq!(scale.band_width(), 30.);
155    }
156
157    #[test]
158    fn test_scale_band_step() {
159        // Adjacent bands start one step apart, whatever the padding.
160        let scale = ScaleBand::new(vec![1, 2, 3], vec![0., 90.]);
161        assert_eq!(
162            scale.step(),
163            scale.tick(&2).unwrap() - scale.tick(&1).unwrap()
164        );
165
166        let padded = ScaleBand::new(vec![1, 2, 3], vec![0., 90.])
167            .padding_inner(0.4)
168            .padding_outer(0.2);
169        assert!(
170            (padded.step() - (padded.tick(&2).unwrap() - padded.tick(&1).unwrap())).abs() < 1e-4
171        );
172
173        // A single band spans the range.
174        assert_eq!(ScaleBand::new(vec![1], vec![0., 90.]).step(), 90.);
175    }
176
177    #[test]
178    fn test_scale_band_zero() {
179        let scale = ScaleBand::new(vec![], vec![0., 90.]);
180        assert_eq!(scale.tick(&1), None);
181        assert_eq!(scale.tick(&2), None);
182        assert_eq!(scale.tick(&3), None);
183        assert_eq!(scale.band_width(), 0.);
184
185        let scale = ScaleBand::new(vec![1, 2, 3], vec![]);
186        assert_eq!(scale.tick(&1), Some(0.));
187        assert_eq!(scale.tick(&2), Some(0.));
188        assert_eq!(scale.tick(&3), Some(0.));
189        assert_eq!(scale.band_width(), 0.);
190    }
191}