1use std::{hash::Hash, 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, PlotAppear, PlotAxis,
13 scale::{PlotValue, Scale, ScaleLinear, ScalePoint},
14 shape::Area,
15 tooltip::{CrossLine, Dot, Tooltip, TooltipState},
16 },
17};
18
19use super::{
20 AXIS_GAP, ChartAppear, 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, reveal_mask,
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 appear: ChartAppear,
49}
50
51impl<T, X, Y> AreaChart<T, X, Y>
52where
53 X: Clone + PartialEq + Into<SharedString> + 'static,
54 Y: PlotValue,
55{
56 #[track_caller]
57 pub fn new<I>(data: I) -> Self
58 where
59 I: IntoIterator<Item = T>,
60 {
61 Self {
62 data: data.into_iter().collect(),
63 curves: vec![],
64 strokes: vec![],
65 fills: vec![],
66 names: vec![],
67 tooltip_content: TooltipContent::default(),
68 tick_margin: 1,
69 x: None,
70 y: vec![],
71 x_axis: true,
72 grid: true,
73 y_domain: None,
74 point_count: None,
75 axes: PointAxes::default(),
76 id: caller_id(),
77 interactive: true,
78 appear: ChartAppear::default(),
79 }
80 }
81
82 pub fn id(mut self, id: impl Into<ElementId>) -> Self {
89 self.id = id.into();
90 self
91 }
92
93 pub fn interactive(mut self, interactive: bool) -> Self {
101 self.interactive = interactive;
102 self
103 }
104
105 pub fn appear(mut self, appear: bool) -> Self {
112 self.appear.set_enabled(appear);
113 self
114 }
115
116 pub fn appear_key(mut self, key: impl Hash) -> Self {
121 self.appear.set_key(key);
122 self
123 }
124
125 pub fn name(mut self, name: impl Into<SharedString>) -> Self {
129 self.names.push(name.into());
130 self
131 }
132
133 pub fn tooltip_title(mut self, title: impl Fn(&T) -> SharedString + 'static) -> Self {
135 self.tooltip_content.set_title(title);
136 self
137 }
138
139 pub fn tooltip_value(
144 mut self,
145 value: impl Fn(&T, usize, f64) -> SharedString + 'static,
146 ) -> Self {
147 self.tooltip_content.set_value(value);
148 self
149 }
150
151 pub fn tooltip_value_color<H>(mut self, color: impl Fn(&T, usize, f64) -> H + 'static) -> Self
157 where
158 H: Into<Hsla>,
159 {
160 self.tooltip_content.set_value_color(color);
161 self
162 }
163
164 pub fn tooltip_content<E>(
171 mut self,
172 content: impl Fn(&T, &mut Window, &mut App) -> E + 'static,
173 ) -> Self
174 where
175 E: IntoElement,
176 {
177 self.tooltip_content.set_content(content);
178 self
179 }
180
181 pub fn x(mut self, x: impl Fn(&T) -> X + 'static) -> Self {
182 self.x = Some(Rc::new(x));
183 self
184 }
185
186 pub fn y(mut self, y: impl Fn(&T) -> Y + 'static) -> Self {
187 self.y.push(Rc::new(y));
188 self
189 }
190
191 pub fn stroke(mut self, stroke: impl Into<Hsla>) -> Self {
192 self.strokes.push(stroke.into());
193 self
194 }
195
196 pub fn fill(mut self, fill: impl Into<Background>) -> Self {
197 self.fills.push(fill.into());
198 self
199 }
200
201 pub fn natural(mut self) -> Self {
202 self.curves.push(Curve::Natural);
203 self
204 }
205
206 pub fn linear(mut self) -> Self {
207 self.curves.push(Curve::Linear);
208 self
209 }
210
211 pub fn step_after(mut self) -> Self {
212 self.curves.push(Curve::StepAfter);
213 self
214 }
215
216 pub fn tick_margin(mut self, tick_margin: usize) -> Self {
217 self.tick_margin = tick_margin;
218 self
219 }
220
221 pub fn x_axis(mut self, x_axis: bool) -> Self {
225 self.x_axis = x_axis;
226 self
227 }
228
229 pub fn grid(mut self, grid: bool) -> Self {
230 self.grid = grid;
231 self
232 }
233
234 pub fn y_domain(mut self, min: Y, max: Y) -> Self {
242 self.y_domain = Some((min, max));
243 self
244 }
245
246 pub fn point_count(mut self, count: usize) -> Self {
255 self.point_count = Some(count);
256 self
257 }
258
259 pub fn y_axis(mut self, y_axis: bool) -> Self {
263 self.axes.y_axis = y_axis;
264 self
265 }
266
267 pub fn y_axis_label_placement(mut self, placement: AxisLabelPlacement) -> Self {
272 self.axes.y_axis_label_placement = placement;
273 self
274 }
275
276 pub fn y_tick_count(mut self, count: usize) -> Self {
285 self.axes.y_tick_count = count.max(2);
286 self
287 }
288
289 pub fn y_tick_format<S>(mut self, format: impl Fn(f64) -> S + 'static) -> Self
291 where
292 S: Into<SharedString> + 'static,
293 {
294 self.axes.y_tick_format = Some(Rc::new(move |value| format(value).into()));
295 self
296 }
297
298 pub fn x_tick_count(mut self, count: usize) -> Self {
305 self.axes.x_tick_count = Some(count);
306 self
307 }
308
309 pub fn grid_columns(mut self, count: usize) -> Self {
314 self.axes.grid_columns = count;
315 self
316 }
317
318 pub fn grid_dashed(mut self, dashed: bool) -> Self {
322 self.axes.grid_dashed = dashed;
323 self
324 }
325
326 pub fn reference_line(mut self, value: Y) -> Self {
330 if let Some(value) = value.to_f64() {
331 self.axes.reference_lines.push(value);
332 }
333 self
334 }
335
336 pub fn y_padding(mut self, top: f32, bottom: f32) -> Self {
341 self.axes.y_padding = (top, bottom);
342 self
343 }
344
345 fn scales(
350 &self,
351 bounds: Bounds<Pixels>,
352 ) -> Option<(ScalePoint<X>, ScaleLinear<Y>, ValueExtent)> {
353 let x_fn = self.x.as_ref()?;
354 if self.y.is_empty() {
355 return None;
356 }
357
358 let width = bounds.size.width.as_f32();
359 let axis_gap = if self.x_axis { AXIS_GAP } else { 0. };
360 let height = bounds.size.height.as_f32() - axis_gap;
361
362 let len = self.data.len();
363 let x = ScalePoint::new(
364 self.data.iter().map(|v| x_fn(v)),
365 point_range(
366 self.axes.plot_left(),
367 width - self.axes.plot_left(),
368 len,
369 axis_point_count(self.point_count, len),
370 ),
371 );
372 let (y, extent) = point_value_scale(
373 self.data
374 .iter()
375 .flat_map(|v| self.y.iter().map(|y_fn| y_fn(v))),
376 self.y_domain,
377 height,
378 self.axes.y_padding,
379 );
380
381 Some((x, y, extent))
382 }
383}
384
385impl<T, X, Y> Plot for AreaChart<T, X, Y>
386where
387 X: Clone + PartialEq + Into<SharedString> + 'static,
388 Y: PlotValue,
389{
390 fn prepaint(
391 &mut self,
392 bounds: Bounds<Pixels>,
393 window: &mut Window,
394 _cx: &mut App,
395 ) -> Vec<AnyElement> {
396 if let Some((_, _, extent)) = self.scales(bounds) {
398 let axis_gap = if self.x_axis { AXIS_GAP } else { 0. };
399 let height = bounds.size.height.as_f32() - axis_gap;
400 self.axes.measure_y_labels(extent, height, window);
401 }
402 vec![]
403 }
404
405 fn paint(&mut self, bounds: Bounds<Pixels>, window: &mut Window, cx: &mut App) {
406 let Some(x_fn) = self.x.as_ref() else {
407 return;
408 };
409 let Some((x, y, extent)) = self.scales(bounds) else {
410 return;
411 };
412
413 let axis_gap = if self.x_axis { AXIS_GAP } else { 0. };
414 let height = bounds.size.height.as_f32() - axis_gap;
415
416 let left = self.axes.plot_left();
420 let axis_bounds = Bounds {
421 origin: bounds.origin + point(px(left), px(0.)),
422 size: Size::new(bounds.size.width - px(left), bounds.size.height),
423 };
424 let mut axis = PlotAxis::new().stroke(cx.theme().border);
425 if self.x_axis {
426 let labeled = labeled_items(
427 axis_point_count(self.point_count, self.data.len()),
428 self.axes.x_tick_count,
429 self.tick_margin,
430 );
431 let labels = build_point_x_labels(
432 &self.data,
433 x_fn.as_ref(),
434 &x,
435 axis_point_count(self.point_count, self.data.len()),
436 &labeled,
437 cx.theme().muted_foreground,
438 )
439 .into_iter()
440 .map(|mut label| {
441 label.tick -= px(left);
442 label
443 });
444 axis = axis.x(height).x_label(labels);
445 }
446 axis.paint(&axis_bounds, window, cx);
447
448 if self.grid {
449 self.axes.paint_grid(bounds, height, window, cx);
450 }
451
452 let default_fill: Background = cx.theme().chart_2.opacity(0.4).into();
454 let default_stroke = cx.theme().chart_2;
455 let areas = self.y.iter().enumerate().map(|(i, y_fn)| {
456 let x = x.clone();
457 let y = y.clone();
458 let y_fn = y_fn.clone();
459
460 let fill = *self.fills.get(i).unwrap_or(&default_fill);
461 let stroke = *self.strokes.get(i).unwrap_or(&default_stroke);
462 let curve = *self
463 .curves
464 .get(i)
465 .unwrap_or(self.curves.first().unwrap_or(&Default::default()));
466
467 Area::new()
468 .data(self.data.iter().enumerate())
470 .x(move |(i, _)| x.tick_at(*i))
471 .y0(height)
472 .y1(move |(_, d)| y.tick(&y_fn(d)))
473 .stroke(stroke)
474 .curve(curve)
475 .fill(fill)
476 });
477
478 let mask = self
479 .y_domain
480 .is_some()
481 .then(|| pinned_plot_mask(bounds, height));
482 let reveal = reveal_mask(bounds, left, self.appear.get().progress());
485 window.with_content_mask(mask, |window| {
486 window.with_content_mask(reveal, |window| {
487 let caches = PathCaches::for_paint("areas", window, cx);
488 caches.update(cx, |caches, _| {
489 for (i, area) in areas.enumerate() {
490 let (fill, line) = caches.slot_pair(i);
491 area.paint_cached(&bounds, fill, line, window);
492 }
493 });
494 });
495 });
496
497 self.axes
498 .paint_reference_lines(extent, bounds, height, window, cx);
499 self.axes.paint_y_labels(extent, bounds, height, window, cx);
500 }
501
502 fn id(&self) -> Option<ElementId> {
503 Some(self.id.clone())
504 }
505
506 fn interactive(&self) -> bool {
507 self.interactive
508 }
509
510 fn appear(&mut self, appear: PlotAppear, _window: &mut Window, _cx: &mut App) {
511 self.appear.update(appear);
512 }
513
514 fn appear_generation(&self) -> Option<u64> {
515 self.appear.generation()
516 }
517
518 fn tooltip_state(
519 &self,
520 position: Point<Pixels>,
521 bounds: Bounds<Pixels>,
522 _cx: &App,
523 ) -> Option<TooltipState> {
524 let (x, y, _) = self.scales(bounds)?;
525
526 let axis_gap = if self.x_axis { AXIS_GAP } else { 0. };
528 if position.y.as_f32() > bounds.size.height.as_f32() - axis_gap
529 || position.x.as_f32() < self.axes.plot_left()
530 {
531 return None;
532 }
533
534 let index = x.nearest_index(position.x.as_f32());
535 let d = self.data.get(index)?;
536 let x_tick = x.tick_at(index)?;
537
538 let dots = self
540 .y
541 .iter()
542 .filter_map(|y_fn| Some(point(px(x_tick), px(y.tick(&y_fn(d))?))))
543 .collect();
544
545 Some(TooltipState::new(
546 index,
547 point(px(x_tick), position.y),
548 dots,
549 ))
550 }
551
552 fn tooltip(
553 &self,
554 state: &TooltipState,
555 cursor: Point<Pixels>,
556 bounds: Bounds<Pixels>,
557 window: &mut Window,
558 cx: &mut App,
559 ) -> Option<AnyElement> {
560 let x_fn = self.x.as_ref()?;
561 let d = self.data.get(state.index)?;
562
563 let default_color = cx.theme().chart_2;
564 let dot_stroke = cx.theme().background;
565 let color = |i: usize| *self.strokes.get(i).unwrap_or(&default_color);
566
567 let tooltip = Tooltip::new(cursor, bounds.size)
569 .gap(px(8.))
570 .cross_line(
572 CrossLine::new(state.cross_line)
573 .height(bounds.size.height.as_f32() - if self.x_axis { AXIS_GAP } else { 0. }),
574 )
575 .dots(state.dots.iter().enumerate().map(|(i, p)| {
576 Dot::new(*p)
577 .size(HOVER_DOT_SIZE)
578 .halo(HOVER_HALO_SIZE)
579 .stroke(dot_stroke)
580 .fill(color(i))
581 }));
582
583 let tooltip = self.tooltip_content.apply(
584 tooltip,
585 d,
586 || Some(x_fn(d).into()),
587 || {
589 self.y
590 .iter()
591 .enumerate()
592 .map(|(i, y_fn)| {
593 let name = self.names.get(i).cloned().unwrap_or_default();
594 Some((color(i), name, y_fn(d).to_f64()?))
595 })
596 .collect::<Option<Vec<_>>>()
597 },
598 window,
599 cx,
600 )?;
601
602 Some(tooltip.into_any_element())
603 }
604}
605
606#[cfg(test)]
607mod tests {
608 use gpui::{Bounds, point, px, size};
609
610 use super::AreaChart;
611 use crate::plot::scale::Scale;
612
613 fn bounds() -> Bounds<gpui::Pixels> {
614 Bounds::new(point(px(0.), px(0.)), size(px(100.), px(50.)))
615 }
616
617 fn chart(data: Vec<f64>) -> AreaChart<(usize, f64), String, f64> {
618 AreaChart::new(data.into_iter().enumerate())
619 .x(|(i, _)| i.to_string())
620 .y(|(_, v)| *v)
621 .x_axis(false)
622 }
623
624 #[test]
625 fn test_point_count_fills_the_leading_part() {
626 let (x, _, _) = chart(vec![1., 2., 3.])
627 .point_count(5)
628 .scales(bounds())
629 .unwrap();
630 assert_eq!(x.tick(&"0".to_string()), Some(0.));
631 assert_eq!(x.tick(&"2".to_string()), Some(50.));
632
633 let (x, _, _) = chart(vec![1., 2., 3.])
634 .point_count(2)
635 .scales(bounds())
636 .unwrap();
637 assert_eq!(x.tick(&"2".to_string()), Some(100.));
638 }
639
640 #[test]
641 fn test_y_domain_replaces_the_fit_from_zero() {
642 let (_, y, _) = chart(vec![10., 20.])
643 .y_domain(10., 20.)
644 .scales(bounds())
645 .unwrap();
646 assert_eq!(y.tick(&10.), Some(50.));
647 assert_eq!(y.tick(&20.), Some(10.));
648
649 let (_, y, _) = chart(vec![10., 20.]).scales(bounds()).unwrap();
650 assert_eq!(y.tick(&0.), Some(50.));
651 assert_eq!(y.tick(&20.), Some(10.));
652 }
653}