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 pub fn id(mut self, id: impl Into<ElementId>) -> Self {
87 self.id = id.into();
88 self
89 }
90
91 pub fn interactive(mut self, interactive: bool) -> Self {
100 self.interactive = interactive;
101 self
102 }
103
104 pub fn name(mut self, name: impl Into<SharedString>) -> Self {
108 self.names.push(name.into());
109 self
110 }
111
112 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 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 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 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 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 pub fn y_domain(mut self, min: Y, max: Y) -> Self {
221 self.y_domain = Some((min, max));
222 self
223 }
224
225 pub fn point_count(mut self, count: usize) -> Self {
234 self.point_count = Some(count);
235 self
236 }
237
238 pub fn y_axis(mut self, y_axis: bool) -> Self {
242 self.axes.y_axis = y_axis;
243 self
244 }
245
246 pub fn y_axis_label_placement(mut self, placement: AxisLabelPlacement) -> Self {
251 self.axes.y_axis_label_placement = placement;
252 self
253 }
254
255 pub fn y_tick_count(mut self, count: usize) -> Self {
264 self.axes.y_tick_count = count.max(2);
265 self
266 }
267
268 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 pub fn x_tick_count(mut self, count: usize) -> Self {
284 self.axes.x_tick_count = Some(count);
285 self
286 }
287
288 pub fn grid_columns(mut self, count: usize) -> Self {
293 self.axes.grid_columns = count;
294 self
295 }
296
297 pub fn grid_dashed(mut self, dashed: bool) -> Self {
301 self.axes.grid_dashed = dashed;
302 self
303 }
304
305 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 pub fn y_padding(mut self, top: f32, bottom: f32) -> Self {
320 self.axes.y_padding = (top, bottom);
321 self
322 }
323
324 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 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 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 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 .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 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 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 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 let tooltip = Tooltip::new(cursor, bounds.size)
540 .gap(px(8.))
541 .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 || {
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}