Skip to main content

gpui_component/chart/
area_chart.rs

1use std::rc::Rc;
2
3use gpui::{
4    AnyElement, App, Background, Bounds, ElementId, Hsla, IntoElement, Pixels, Point, SharedString,
5    Size, Window, point, px,
6};
7use gpui_component_macros::IntoPlot;
8
9use crate::{
10    ActiveTheme,
11    plot::{
12        AxisLabelPlacement, Curve, PathCaches, Plot, PlotAxis,
13        scale::{PlotValue, Scale, ScaleLinear, ScalePoint},
14        shape::Area,
15        tooltip::{CrossLine, Dot, Tooltip, TooltipState},
16    },
17};
18
19use super::{
20    AXIS_GAP, HOVER_DOT_SIZE, HOVER_HALO_SIZE, PointAxes, TooltipContent, ValueExtent,
21    axis_point_count, build_point_x_labels, caller_id, labeled_items, pinned_plot_mask,
22    point_range, point_value_scale,
23};
24
25#[derive(IntoPlot)]
26pub struct AreaChart<T, X, Y>
27where
28    T: 'static,
29    X: Clone + PartialEq + Into<SharedString> + 'static,
30    Y: PlotValue,
31{
32    data: Vec<T>,
33    x: Option<Rc<dyn Fn(&T) -> X>>,
34    y: Vec<Rc<dyn Fn(&T) -> Y>>,
35    strokes: Vec<Hsla>,
36    curves: Vec<Curve>,
37    fills: Vec<Background>,
38    names: Vec<SharedString>,
39    tooltip_content: TooltipContent<T>,
40    tick_margin: usize,
41    x_axis: bool,
42    grid: bool,
43    y_domain: Option<(Y, Y)>,
44    point_count: Option<usize>,
45    axes: PointAxes,
46    id: ElementId,
47    interactive: bool,
48}
49
50impl<T, X, Y> AreaChart<T, X, Y>
51where
52    X: Clone + PartialEq + Into<SharedString> + 'static,
53    Y: PlotValue,
54{
55    #[track_caller]
56    pub fn new<I>(data: I) -> Self
57    where
58        I: IntoIterator<Item = T>,
59    {
60        Self {
61            data: data.into_iter().collect(),
62            curves: vec![],
63            strokes: vec![],
64            fills: vec![],
65            names: vec![],
66            tooltip_content: TooltipContent::default(),
67            tick_margin: 1,
68            x: None,
69            y: vec![],
70            x_axis: true,
71            grid: true,
72            y_domain: None,
73            point_count: None,
74            axes: PointAxes::default(),
75            id: caller_id(),
76            interactive: true,
77        }
78    }
79
80    /// Name this chart's [`ElementId`], replacing the default taken from the
81    /// construction site.
82    ///
83    /// Pass one where a single construction site renders several of these
84    /// charts as siblings: they share the default id, and with it one hover
85    /// state and one path cache. The id must be unique among those siblings.
86    pub fn id(mut self, id: impl Into<ElementId>) -> Self {
87        self.id = id.into();
88        self
89    }
90
91    /// Turn this chart's interactive layer on or off. On by default.
92    ///
93    /// The layer is the hitbox under the cursor and what it drives: a crosshair
94    /// and a dot per series mark the hovered point, and a tooltip shows a row
95    /// each. Turn it off for a chart that only decorates, or one an element above
96    /// it wants the cursor for: without a hitbox it neither answers the mouse nor
97    /// takes the hover from what sits over it. A chart that is off also drops its
98    /// path cache, which is keyed on the same id.
99    pub fn interactive(mut self, interactive: bool) -> Self {
100        self.interactive = interactive;
101        self
102    }
103
104    /// Set the name of the most recently added series, shown in its tooltip row.
105    ///
106    /// Call after the matching [`AreaChart::y`] (e.g. `.y(..).stroke(..).name("Desktop")`).
107    pub fn name(mut self, name: impl Into<SharedString>) -> Self {
108        self.names.push(name.into());
109        self
110    }
111
112    /// Set the hover tooltip's title for a datum, instead of its x value.
113    pub fn tooltip_title(mut self, title: impl Fn(&T) -> SharedString + 'static) -> Self {
114        self.tooltip_content.set_title(title);
115        self
116    }
117
118    /// Set the text of each tooltip row's value; the raw number by default.
119    ///
120    /// The closure receives the datum, the row's index (the series' index in the order `y` added
121    /// them) and the value the row reads.
122    pub fn tooltip_value(
123        mut self,
124        value: impl Fn(&T, usize, f64) -> SharedString + 'static,
125    ) -> Self {
126        self.tooltip_content.set_value(value);
127        self
128    }
129
130    /// Color each tooltip row's value, such as green or red by its sign; the
131    /// tooltip's text color by default.
132    ///
133    /// The closure receives the same arguments as
134    /// [`tooltip_value`](Self::tooltip_value).
135    pub fn tooltip_value_color<H>(mut self, color: impl Fn(&T, usize, f64) -> H + 'static) -> Self
136    where
137        H: Into<Hsla>,
138    {
139        self.tooltip_content.set_value_color(color);
140        self
141    }
142
143    /// Draw the tooltip box's content for a datum yourself, in place of the
144    /// title and rows, for a layout they cannot express such as a table.
145    ///
146    /// The crosshair, the dots and where the box sits stay the chart's, and
147    /// [`tooltip_title`](Self::tooltip_title), [`tooltip_value`](Self::tooltip_value)
148    /// and [`tooltip_value_color`](Self::tooltip_value_color) no longer apply.
149    pub fn tooltip_content<E>(
150        mut self,
151        content: impl Fn(&T, &mut Window, &mut App) -> E + 'static,
152    ) -> Self
153    where
154        E: IntoElement,
155    {
156        self.tooltip_content.set_content(content);
157        self
158    }
159
160    pub fn x(mut self, x: impl Fn(&T) -> X + 'static) -> Self {
161        self.x = Some(Rc::new(x));
162        self
163    }
164
165    pub fn y(mut self, y: impl Fn(&T) -> Y + 'static) -> Self {
166        self.y.push(Rc::new(y));
167        self
168    }
169
170    pub fn stroke(mut self, stroke: impl Into<Hsla>) -> Self {
171        self.strokes.push(stroke.into());
172        self
173    }
174
175    pub fn fill(mut self, fill: impl Into<Background>) -> Self {
176        self.fills.push(fill.into());
177        self
178    }
179
180    pub fn natural(mut self) -> Self {
181        self.curves.push(Curve::Natural);
182        self
183    }
184
185    pub fn linear(mut self) -> Self {
186        self.curves.push(Curve::Linear);
187        self
188    }
189
190    pub fn step_after(mut self) -> Self {
191        self.curves.push(Curve::StepAfter);
192        self
193    }
194
195    pub fn tick_margin(mut self, tick_margin: usize) -> Self {
196        self.tick_margin = tick_margin;
197        self
198    }
199
200    /// Show or hide the x-axis line and labels.
201    ///
202    /// Default is true.
203    pub fn x_axis(mut self, x_axis: bool) -> Self {
204        self.x_axis = x_axis;
205        self
206    }
207
208    pub fn grid(mut self, grid: bool) -> Self {
209        self.grid = grid;
210        self
211    }
212
213    /// Pin the y axis to `min..=max` instead of fitting every series from zero.
214    ///
215    /// Pin it where zero is not a meaningful baseline, such as a price line
216    /// that would otherwise be pressed flat against the top. The range keeps
217    /// the `y_padding` headroom above `max`, 10px by default,
218    /// and the series are clipped to the plot, so a value outside the range
219    /// stops at its edge. Nothing is drawn when `min` equals `max`.
220    pub fn y_domain(mut self, min: Y, max: Y) -> Self {
221        self.y_domain = Some((min, max));
222        self
223    }
224
225    /// Lay the x axis out for `count` evenly spaced points instead of the
226    /// data's own length.
227    ///
228    /// The data takes the leading points in order, the i-th item on the i-th
229    /// point, and the rest stay empty, as an intraday chart does before the
230    /// close. The data has to be contiguous from the first point: a missing
231    /// item shifts every later one a point to the left. A `count` below the
232    /// data's length has no effect.
233    pub fn point_count(mut self, count: usize) -> Self {
234        self.point_count = Some(count);
235        self
236    }
237
238    /// Show the y axis's tick labels, one at each of the `y_tick_count` ticks.
239    ///
240    /// Default is false.
241    pub fn y_axis(mut self, y_axis: bool) -> Self {
242        self.axes.y_axis = y_axis;
243        self
244    }
245
246    /// Set where the y-axis tick labels sit: in a gutter left of the plot, or
247    /// inside it beside their grid lines.
248    ///
249    /// Default is [`AxisLabelPlacement::Outside`].
250    pub fn y_axis_label_placement(mut self, placement: AxisLabelPlacement) -> Self {
251        self.axes.y_axis_label_placement = placement;
252        self
253    }
254
255    /// Set how many ticks the y axis carries, evenly spaced from the baseline
256    /// to the top edge with both ends included.
257    ///
258    /// The ticks place the horizontal grid lines and the tick labels, and each
259    /// label reads the value the scale puts at its height. Values below 2 are
260    /// raised to 2.
261    ///
262    /// Default is 5.
263    pub fn y_tick_count(mut self, count: usize) -> Self {
264        self.axes.y_tick_count = count.max(2);
265        self
266    }
267
268    /// Set the text of each y-axis tick label from the value at its tick.
269    pub fn y_tick_format<S>(mut self, format: impl Fn(f64) -> S + 'static) -> Self
270    where
271        S: Into<SharedString> + 'static,
272    {
273        self.axes.y_tick_format = Some(Rc::new(move |value| format(value).into()));
274        self
275    }
276
277    /// Label `count` of the x values, spread evenly from the first to the
278    /// last, instead of every `tick_margin`-th.
279    ///
280    /// With [`Self::point_count`] set, the labels spread over all the points the
281    /// axis is laid out for, so they keep their places as the data grows; one
282    /// that falls past the data is not drawn yet.
283    pub fn x_tick_count(mut self, count: usize) -> Self {
284        self.axes.x_tick_count = Some(count);
285        self
286    }
287
288    /// Divide the plot into `count` columns with vertical grid lines, the first
289    /// on its left edge.
290    ///
291    /// Default is 0, no vertical lines.
292    pub fn grid_columns(mut self, count: usize) -> Self {
293        self.axes.grid_columns = count;
294        self
295    }
296
297    /// Draw the grid dashed or solid.
298    ///
299    /// Default is true.
300    pub fn grid_dashed(mut self, dashed: bool) -> Self {
301        self.axes.grid_dashed = dashed;
302        self
303    }
304
305    /// Draw a dashed line across the plot at `value`, such as a previous close.
306    ///
307    /// Call again for more lines. A value outside the y axis is not drawn.
308    pub fn reference_line(mut self, value: Y) -> Self {
309        if let Some(value) = value.to_f64() {
310            self.axes.reference_lines.push(value);
311        }
312        self
313    }
314
315    /// Set the space kept clear above the highest value and below the lowest,
316    /// in pixels.
317    ///
318    /// Default is 10px above and none below.
319    pub fn y_padding(mut self, top: f32, bottom: f32) -> Self {
320        self.axes.y_padding = (top, bottom);
321        self
322    }
323
324    /// Build the x (point) and y (linear) scales for the given bounds.
325    ///
326    /// Shared by `paint` and `tooltip_state` so the two stay in sync. Returns `None` when there
327    /// is no x accessor or no series.
328    fn scales(
329        &self,
330        bounds: Bounds<Pixels>,
331    ) -> Option<(ScalePoint<X>, ScaleLinear<Y>, ValueExtent)> {
332        let x_fn = self.x.as_ref()?;
333        if self.y.is_empty() {
334            return None;
335        }
336
337        let width = bounds.size.width.as_f32();
338        let axis_gap = if self.x_axis { AXIS_GAP } else { 0. };
339        let height = bounds.size.height.as_f32() - axis_gap;
340
341        let len = self.data.len();
342        let x = ScalePoint::new(
343            self.data.iter().map(|v| x_fn(v)),
344            point_range(
345                self.axes.plot_left(),
346                width - self.axes.plot_left(),
347                len,
348                axis_point_count(self.point_count, len),
349            ),
350        );
351        let (y, extent) = point_value_scale(
352            self.data
353                .iter()
354                .flat_map(|v| self.y.iter().map(|y_fn| y_fn(v))),
355            self.y_domain,
356            height,
357            self.axes.y_padding,
358        );
359
360        Some((x, y, extent))
361    }
362}
363
364impl<T, X, Y> Plot for AreaChart<T, X, Y>
365where
366    X: Clone + PartialEq + Into<SharedString> + 'static,
367    Y: PlotValue,
368{
369    fn prepaint(
370        &mut self,
371        bounds: Bounds<Pixels>,
372        window: &mut Window,
373        _cx: &mut App,
374    ) -> Vec<AnyElement> {
375        // The y labels' gutter is measured before the x scale is laid out past it.
376        if let Some((_, _, extent)) = self.scales(bounds) {
377            let axis_gap = if self.x_axis { AXIS_GAP } else { 0. };
378            let height = bounds.size.height.as_f32() - axis_gap;
379            self.axes.measure_y_labels(extent, height, window);
380        }
381        vec![]
382    }
383
384    fn paint(&mut self, bounds: Bounds<Pixels>, window: &mut Window, cx: &mut App) {
385        let Some(x_fn) = self.x.as_ref() else {
386            return;
387        };
388        let Some((x, y, extent)) = self.scales(bounds) else {
389            return;
390        };
391
392        let axis_gap = if self.x_axis { AXIS_GAP } else { 0. };
393        let height = bounds.size.height.as_f32() - axis_gap;
394
395        // Draw X axis
396        // The axis runs under the plot only, clear of a value-axis gutter, so
397        // its labels shift back by the gutter the x scale already includes.
398        let left = self.axes.plot_left();
399        let axis_bounds = Bounds {
400            origin: bounds.origin + point(px(left), px(0.)),
401            size: Size::new(bounds.size.width - px(left), bounds.size.height),
402        };
403        let mut axis = PlotAxis::new().stroke(cx.theme().border);
404        if self.x_axis {
405            let labeled = labeled_items(
406                axis_point_count(self.point_count, self.data.len()),
407                self.axes.x_tick_count,
408                self.tick_margin,
409            );
410            let labels = build_point_x_labels(
411                &self.data,
412                x_fn.as_ref(),
413                &x,
414                axis_point_count(self.point_count, self.data.len()),
415                &labeled,
416                cx.theme().muted_foreground,
417            )
418            .into_iter()
419            .map(|mut label| {
420                label.tick -= px(left);
421                label
422            });
423            axis = axis.x(height).x_label(labels);
424        }
425        axis.paint(&axis_bounds, window, cx);
426
427        if self.grid {
428            self.axes.paint_grid(bounds, height, window, cx);
429        }
430
431        // Draw area
432        let default_fill: Background = cx.theme().chart_2.opacity(0.4).into();
433        let default_stroke = cx.theme().chart_2;
434        let areas = self.y.iter().enumerate().map(|(i, y_fn)| {
435            let x = x.clone();
436            let y = y.clone();
437            let y_fn = y_fn.clone();
438
439            let fill = *self.fills.get(i).unwrap_or(&default_fill);
440            let stroke = *self.strokes.get(i).unwrap_or(&default_stroke);
441            let curve = *self
442                .curves
443                .get(i)
444                .unwrap_or(self.curves.first().unwrap_or(&Default::default()));
445
446            Area::new()
447                // One x domain entry per datum: project by index, not by lookup.
448                .data(self.data.iter().enumerate())
449                .x(move |(i, _)| x.tick_at(*i))
450                .y0(height)
451                .y1(move |(_, d)| y.tick(&y_fn(d)))
452                .stroke(stroke)
453                .curve(curve)
454                .fill(fill)
455        });
456
457        let mask = self
458            .y_domain
459            .is_some()
460            .then(|| pinned_plot_mask(bounds, height));
461        window.with_content_mask(mask, |window| {
462            // Caching hangs off the chart's own id, which only an interactive chart
463            // puts on the stack; without one, siblings would share a slot and thrash
464            // it, so a chart that is off tessellates afresh each paint.
465            if self.interactive {
466                let caches = PathCaches::for_paint("areas", window, cx);
467                caches.update(cx, |caches, _| {
468                    for (i, area) in areas.enumerate() {
469                        let (fill, line) = caches.slot_pair(i);
470                        area.paint_cached(&bounds, fill, line, window);
471                    }
472                });
473            } else {
474                for area in areas {
475                    area.paint(&bounds, window);
476                }
477            }
478        });
479
480        self.axes
481            .paint_reference_lines(extent, bounds, height, window, cx);
482        self.axes.paint_y_labels(extent, bounds, height, window, cx);
483    }
484
485    fn id(&self) -> Option<ElementId> {
486        self.interactive.then(|| self.id.clone())
487    }
488
489    fn tooltip_state(
490        &self,
491        position: Point<Pixels>,
492        bounds: Bounds<Pixels>,
493        _cx: &App,
494    ) -> Option<TooltipState> {
495        let (x, y, _) = self.scales(bounds)?;
496
497        // Ignore the x-axis label gutter so hovering the labels doesn't show a tooltip.
498        let axis_gap = if self.x_axis { AXIS_GAP } else { 0. };
499        if position.y.as_f32() > bounds.size.height.as_f32() - axis_gap
500            || position.x.as_f32() < self.axes.plot_left()
501        {
502            return None;
503        }
504
505        let index = x.nearest_index(position.x.as_f32());
506        let d = self.data.get(index)?;
507        let x_tick = x.tick_at(index)?;
508
509        // One dot per series at the hovered x.
510        let dots = self
511            .y
512            .iter()
513            .filter_map(|y_fn| Some(point(px(x_tick), px(y.tick(&y_fn(d))?))))
514            .collect();
515
516        Some(TooltipState::new(
517            index,
518            point(px(x_tick), position.y),
519            dots,
520        ))
521    }
522
523    fn tooltip(
524        &self,
525        state: &TooltipState,
526        cursor: Point<Pixels>,
527        bounds: Bounds<Pixels>,
528        window: &mut Window,
529        cx: &mut App,
530    ) -> Option<AnyElement> {
531        let x_fn = self.x.as_ref()?;
532        let d = self.data.get(state.index)?;
533
534        let default_color = cx.theme().chart_2;
535        let dot_stroke = cx.theme().background;
536        let color = |i: usize| *self.strokes.get(i).unwrap_or(&default_color);
537
538        // Follow the cursor; the crosshair and dots glide to the data point.
539        let tooltip = Tooltip::new(cursor, bounds.size)
540            .gap(px(8.))
541            // Confine the crosshair to the plot area so it doesn't cross the x-axis.
542            .cross_line(
543                CrossLine::new(state.cross_line)
544                    .height(bounds.size.height.as_f32() - if self.x_axis { AXIS_GAP } else { 0. }),
545            )
546            .dots(state.dots.iter().enumerate().map(|(i, p)| {
547                Dot::new(*p)
548                    .size(HOVER_DOT_SIZE)
549                    .halo(HOVER_HALO_SIZE)
550                    .stroke(dot_stroke)
551                    .fill(color(i))
552            }));
553
554        let tooltip = self.tooltip_content.apply(
555            tooltip,
556            d,
557            || Some(x_fn(d).into()),
558            // One row per series: swatch + label + value.
559            || {
560                self.y
561                    .iter()
562                    .enumerate()
563                    .map(|(i, y_fn)| {
564                        let name = self.names.get(i).cloned().unwrap_or_default();
565                        Some((color(i), name, y_fn(d).to_f64()?))
566                    })
567                    .collect::<Option<Vec<_>>>()
568            },
569            window,
570            cx,
571        )?;
572
573        Some(tooltip.into_any_element())
574    }
575}
576
577#[cfg(test)]
578mod tests {
579    use gpui::{Bounds, point, px, size};
580
581    use super::AreaChart;
582    use crate::plot::scale::Scale;
583
584    fn bounds() -> Bounds<gpui::Pixels> {
585        Bounds::new(point(px(0.), px(0.)), size(px(100.), px(50.)))
586    }
587
588    fn chart(data: Vec<f64>) -> AreaChart<(usize, f64), String, f64> {
589        AreaChart::new(data.into_iter().enumerate())
590            .x(|(i, _)| i.to_string())
591            .y(|(_, v)| *v)
592            .x_axis(false)
593    }
594
595    #[test]
596    fn test_point_count_fills_the_leading_part() {
597        let (x, _, _) = chart(vec![1., 2., 3.])
598            .point_count(5)
599            .scales(bounds())
600            .unwrap();
601        assert_eq!(x.tick(&"0".to_string()), Some(0.));
602        assert_eq!(x.tick(&"2".to_string()), Some(50.));
603
604        let (x, _, _) = chart(vec![1., 2., 3.])
605            .point_count(2)
606            .scales(bounds())
607            .unwrap();
608        assert_eq!(x.tick(&"2".to_string()), Some(100.));
609    }
610
611    #[test]
612    fn test_y_domain_replaces_the_fit_from_zero() {
613        let (_, y, _) = chart(vec![10., 20.])
614            .y_domain(10., 20.)
615            .scales(bounds())
616            .unwrap();
617        assert_eq!(y.tick(&10.), Some(50.));
618        assert_eq!(y.tick(&20.), Some(10.));
619
620        let (_, y, _) = chart(vec![10., 20.]).scales(bounds()).unwrap();
621        assert_eq!(y.tick(&0.), Some(50.));
622        assert_eq!(y.tick(&20.), Some(10.));
623    }
624}