1use std::{
2 f32::consts::{PI, TAU},
3 rc::Rc,
4};
5
6use gpui::{
7 AnyElement, App, AvailableSpace, Background, Bounds, ElementId, Hsla, IntoElement, Pixels,
8 Point, SharedString, TextAlign, Window, point, px,
9};
10use gpui_component_macros::IntoPlot;
11use num_traits::Zero;
12
13use crate::{
14 ActiveTheme,
15 plot::{
16 Plot,
17 label::{PlotLabel, TEXT_SIZE, Text},
18 polygon,
19 scale::{PlotValue, Scale, ScaleLinear},
20 shape::RadialLine,
21 tooltip::{Dot, Tooltip, TooltipState},
22 },
23};
24
25use super::{HOVER_DOT_SIZE, HOVER_HALO_SIZE, TooltipContent, caller_id};
26
27const HALF_PI: f32 = PI / 2.;
28
29const DEFAULT_LABEL_GAP: f32 = 10.;
31
32const DEFAULT_GRID_LEVELS: usize = 4;
34
35pub enum RadarLabel {
37 Text(SharedString),
40 Element(AnyElement),
44}
45
46impl From<&'static str> for RadarLabel {
47 fn from(text: &'static str) -> Self {
48 Self::Text(text.into())
49 }
50}
51
52impl From<String> for RadarLabel {
53 fn from(text: String) -> Self {
54 Self::Text(text.into())
55 }
56}
57
58impl From<SharedString> for RadarLabel {
59 fn from(text: SharedString) -> Self {
60 Self::Text(text)
61 }
62}
63
64impl From<AnyElement> for RadarLabel {
65 fn from(element: AnyElement) -> Self {
66 Self::Element(element)
67 }
68}
69
70#[derive(IntoPlot)]
76pub struct RadarChart<T, Y>
77where
78 T: 'static,
79 Y: PlotValue,
80{
81 data: Vec<T>,
82 values: Vec<Rc<dyn Fn(&T) -> Y>>,
83 strokes: Vec<Hsla>,
84 fills: Vec<Background>,
85 names: Vec<SharedString>,
86 tooltip_content: TooltipContent<T>,
87 label: Option<Rc<dyn Fn(&T) -> RadarLabel + 'static>>,
88 label_texts: Vec<Option<SharedString>>,
92 label_color: Option<Hsla>,
93 label_gap: f32,
94 max_value: Option<Y>,
95 outer_radius: f32,
96 grid: bool,
97 grid_levels: usize,
98 dot: bool,
99 id: ElementId,
100 interactive: bool,
101}
102
103impl<T, Y> RadarChart<T, Y>
104where
105 Y: PlotValue,
106{
107 #[track_caller]
108 pub fn new<I>(data: I) -> Self
109 where
110 I: IntoIterator<Item = T>,
111 {
112 Self {
113 data: data.into_iter().collect(),
114 values: vec![],
115 strokes: vec![],
116 fills: vec![],
117 names: vec![],
118 tooltip_content: TooltipContent::default(),
119 label: None,
120 label_texts: vec![],
121 label_color: None,
122 label_gap: DEFAULT_LABEL_GAP,
123 max_value: None,
124 outer_radius: 0.,
125 grid: true,
126 grid_levels: DEFAULT_GRID_LEVELS,
127 dot: false,
128 id: caller_id(),
129 interactive: true,
130 }
131 }
132
133 pub fn id(mut self, id: impl Into<ElementId>) -> Self {
140 self.id = id.into();
141 self
142 }
143
144 pub fn interactive(mut self, interactive: bool) -> Self {
153 self.interactive = interactive;
154 self
155 }
156
157 pub fn name(mut self, name: impl Into<SharedString>) -> Self {
162 self.names.push(name.into());
163 self
164 }
165
166 pub fn tooltip_title(mut self, title: impl Fn(&T) -> SharedString + 'static) -> Self {
168 self.tooltip_content.set_title(title);
169 self
170 }
171
172 pub fn tooltip_value(
177 mut self,
178 value: impl Fn(&T, usize, f64) -> SharedString + 'static,
179 ) -> Self {
180 self.tooltip_content.set_value(value);
181 self
182 }
183
184 pub fn tooltip_value_color<H>(mut self, color: impl Fn(&T, usize, f64) -> H + 'static) -> Self
190 where
191 H: Into<Hsla>,
192 {
193 self.tooltip_content.set_value_color(color);
194 self
195 }
196
197 pub fn tooltip_content<E>(
204 mut self,
205 content: impl Fn(&T, &mut Window, &mut App) -> E + 'static,
206 ) -> Self
207 where
208 E: IntoElement,
209 {
210 self.tooltip_content.set_content(content);
211 self
212 }
213
214 pub fn value(mut self, value: impl Fn(&T) -> Y + 'static) -> Self {
219 self.values.push(Rc::new(value));
220 self
221 }
222
223 pub fn stroke(mut self, stroke: impl Into<Hsla>) -> Self {
227 self.strokes.push(stroke.into());
228 self
229 }
230
231 pub fn fill(mut self, fill: impl Into<Background>) -> Self {
235 self.fills.push(fill.into());
236 self
237 }
238
239 pub fn label<L>(mut self, label: impl Fn(&T) -> L + 'static) -> Self
256 where
257 L: Into<RadarLabel> + 'static,
258 {
259 self.label = Some(Rc::new(move |d| label(d).into()));
260 self
261 }
262
263 pub fn label_color(mut self, color: impl Into<Hsla>) -> Self {
267 self.label_color = Some(color.into());
268 self
269 }
270
271 pub fn label_gap(mut self, gap: f32) -> Self {
274 self.label_gap = gap;
275 self
276 }
277
278 pub fn max_value(mut self, max_value: Y) -> Self {
282 self.max_value = Some(max_value);
283 self
284 }
285
286 pub fn outer_radius(mut self, outer_radius: f32) -> Self {
290 self.outer_radius = outer_radius;
291 self
292 }
293
294 pub fn grid(mut self, grid: bool) -> Self {
298 self.grid = grid;
299 self
300 }
301
302 pub fn grid_levels(mut self, grid_levels: usize) -> Self {
304 self.grid_levels = grid_levels.max(1);
305 self
306 }
307
308 pub fn dot(mut self) -> Self {
310 self.dot = true;
311 self
312 }
313
314 fn series_stroke(&self, ix: usize, cx: &App) -> Hsla {
318 self.series_stroke_from(&Self::palette(cx), ix)
319 }
320
321 fn palette(cx: &App) -> [Hsla; 5] {
323 [
324 cx.theme().chart_1,
325 cx.theme().chart_2,
326 cx.theme().chart_3,
327 cx.theme().chart_4,
328 cx.theme().chart_5,
329 ]
330 }
331
332 fn series_stroke_from(&self, palette: &[Hsla; 5], ix: usize) -> Hsla {
334 self.strokes
335 .get(ix)
336 .copied()
337 .unwrap_or(palette[ix % palette.len()])
338 }
339
340 fn resolve_outer_radius(&self, bounds: &Bounds<Pixels>) -> f32 {
342 if self.outer_radius.is_zero() {
343 bounds.size.height.as_f32() * 0.4
344 } else {
345 self.outer_radius
346 }
347 }
348
349 fn label_anchor(
356 &self,
357 ix: usize,
358 outer_radius: f32,
359 bounds: &Bounds<Pixels>,
360 ) -> (Point<f32>, Point<f32>) {
361 let label_radius = outer_radius + self.label_gap;
362 let angle = ix as f32 * TAU / self.data.len() as f32 - HALF_PI;
363 let direction = point(angle.cos(), angle.sin());
364
365 let anchor = point(
366 bounds.size.width.as_f32() / 2. + label_radius * direction.x,
367 bounds.size.height.as_f32() / 2. + label_radius * direction.y,
368 );
369
370 (anchor, direction)
371 }
372
373 fn scale(&self, outer_radius: f32) -> ScaleLinear<Y> {
378 let domain = if let Some(max_value) = self.max_value {
379 vec![Y::zero(), max_value]
380 } else {
381 self.data
382 .iter()
383 .flat_map(|d| self.values.iter().map(|value_fn| value_fn(d)))
384 .chain(Some(Y::zero()))
385 .collect()
386 };
387
388 ScaleLinear::new(domain, [0., outer_radius])
389 }
390
391 fn hovered_index(&self, position: Point<Pixels>, bounds: Bounds<Pixels>) -> Option<usize> {
394 let n = self.data.len();
395 if n == 0 {
396 return None;
397 }
398
399 let outer_radius = self.resolve_outer_radius(&bounds);
400 let dx = position.x.as_f32() - bounds.size.width.as_f32() / 2.;
401 let dy = position.y.as_f32() - bounds.size.height.as_f32() / 2.;
402 if dx.hypot(dy) > outer_radius + self.label_gap {
403 return None;
404 }
405
406 let angle = (dy.atan2(dx) + HALF_PI).rem_euclid(TAU);
408 Some((angle * n as f32 / TAU).round() as usize % n)
409 }
410}
411
412impl<T, Y> Plot for RadarChart<T, Y>
413where
414 Y: PlotValue,
415{
416 fn prepaint(
419 &mut self,
420 bounds: Bounds<Pixels>,
421 window: &mut Window,
422 cx: &mut App,
423 ) -> Vec<AnyElement> {
424 self.label_texts.clear();
425
426 let n = self.data.len();
428 if n == 0 || self.values.is_empty() {
429 return vec![];
430 }
431 let Some(label_fn) = self.label.clone() else {
432 return vec![];
433 };
434
435 let outer_radius = self.resolve_outer_radius(&bounds);
436 let mut texts = Vec::with_capacity(n);
437 let mut elements = vec![];
438
439 for (ix, d) in self.data.iter().enumerate() {
440 match label_fn(d) {
441 RadarLabel::Text(text) => texts.push(Some(text)),
442 RadarLabel::Element(mut element) => {
443 texts.push(None);
444
445 let size = element.layout_as_root(AvailableSpace::min_size(), window, cx);
448 let (anchor, direction) = self.label_anchor(ix, outer_radius, &bounds);
449
450 let origin = bounds.origin
455 + point(
456 px(anchor.x + (direction.x - 1.) * size.width.as_f32() / 2.),
457 px(anchor.y + (direction.y - 1.) * size.height.as_f32() / 2.),
458 );
459
460 element.prepaint_at(origin, window, cx);
461 elements.push(element);
462 }
463 }
464 }
465
466 self.label_texts = texts;
467
468 elements
469 }
470
471 fn paint(&mut self, bounds: Bounds<Pixels>, window: &mut Window, cx: &mut App) {
472 let n = self.data.len();
473 if n == 0 || self.values.is_empty() {
474 return;
475 }
476
477 let outer_radius = self.resolve_outer_radius(&bounds);
478 let angle_step = TAU / n as f32;
479 let center_x = bounds.size.width.as_f32() / 2.;
480 let center_y = bounds.size.height.as_f32() / 2.;
481 let scale = self.scale(outer_radius);
482
483 if self.grid {
485 let stroke = cx.theme().chart_grid;
486
487 for level in 1..=self.grid_levels {
488 let radius = outer_radius * level as f32 / self.grid_levels as f32;
489 RadialLine::new()
490 .data(0..n)
491 .angle(move |_, i| Some(i as f32 * angle_step))
492 .radius(move |_, _| Some(radius))
493 .closed()
494 .stroke(stroke)
495 .paint(&bounds, window);
496 }
497
498 for i in 0..n {
499 let angle = i as f32 * angle_step - HALF_PI;
500 let points = [
501 point(center_x, center_y),
502 point(
503 center_x + outer_radius * angle.cos(),
504 center_y + outer_radius * angle.sin(),
505 ),
506 ];
507 if let Some(path) = polygon(&points, &bounds) {
508 window.paint_path(path, stroke);
509 }
510 }
511 }
512
513 for (i, value_fn) in self.values.iter().enumerate() {
515 let stroke = self.series_stroke(i, cx);
516 let fill = self
517 .fills
518 .get(i)
519 .copied()
520 .unwrap_or_else(|| stroke.opacity(0.3).into());
521
522 let scale = scale.clone();
523 let value_fn = value_fn.clone();
524 let mut line = RadialLine::new()
525 .data(&self.data)
526 .angle(move |_, i| Some(i as f32 * angle_step))
527 .radius(move |d, _| scale.tick(&value_fn(d)))
528 .closed()
529 .fill(fill)
530 .stroke(stroke)
531 .stroke_width(2.);
532 if self.dot {
533 line = line.dot().dot_size(8.).dot_fill(stroke);
534 }
535 line.paint(&bounds, window);
536 }
537
538 let label_color = self.label_color.unwrap_or(cx.theme().muted_foreground);
541 let labels = self
542 .label_texts
543 .iter()
544 .enumerate()
545 .filter_map(|(ix, text)| {
546 let text = text.clone()?;
547 let (anchor, direction) = self.label_anchor(ix, outer_radius, &bounds);
548
549 let align = if direction.x > 1e-3 {
553 TextAlign::Left
554 } else if direction.x < -1e-3 {
555 TextAlign::Right
556 } else {
557 TextAlign::Center
558 };
559
560 Some(
561 Text::new(
562 text,
563 point(px(anchor.x), px(anchor.y - TEXT_SIZE / 2.)),
564 label_color,
565 )
566 .align(align),
567 )
568 });
569
570 PlotLabel::new(labels.collect()).paint(&bounds, window, cx);
571 }
572
573 fn id(&self) -> Option<ElementId> {
574 self.interactive.then(|| self.id.clone())
575 }
576
577 fn tooltip_state(
578 &self,
579 position: Point<Pixels>,
580 bounds: Bounds<Pixels>,
581 _cx: &App,
582 ) -> Option<TooltipState> {
583 if self.values.is_empty() {
584 return None;
585 }
586 let index = self.hovered_index(position, bounds)?;
587 let d = self.data.get(index)?;
588
589 let outer_radius = self.resolve_outer_radius(&bounds);
590 let scale = self.scale(outer_radius);
591 let center_x = bounds.size.width.as_f32() / 2.;
592 let center_y = bounds.size.height.as_f32() / 2.;
593 let angle = index as f32 * TAU / self.data.len() as f32 - HALF_PI;
594
595 let dots = self
597 .values
598 .iter()
599 .filter_map(|value_fn| {
600 let radius = scale.tick(&value_fn(d))?;
601 Some(point(
602 px(center_x + radius * angle.cos()),
603 px(center_y + radius * angle.sin()),
604 ))
605 })
606 .collect();
607
608 Some(TooltipState::new(index, position, dots))
609 }
610
611 fn tooltip(
612 &self,
613 state: &TooltipState,
614 cursor: Point<Pixels>,
615 bounds: Bounds<Pixels>,
616 window: &mut Window,
617 cx: &mut App,
618 ) -> Option<AnyElement> {
619 let d = self.data.get(state.index)?;
620
621 let dot_stroke = cx.theme().background;
622
623 let tooltip =
626 Tooltip::new(cursor, bounds.size)
627 .gap(px(8.))
628 .dots(state.dots.iter().enumerate().map(|(i, p)| {
629 Dot::new(*p)
630 .size(HOVER_DOT_SIZE)
631 .halo(HOVER_HALO_SIZE)
632 .stroke(dot_stroke)
633 .fill(self.series_stroke(i, cx))
634 }));
635
636 let palette = Self::palette(cx);
637 let tooltip = self.tooltip_content.apply(
638 tooltip,
639 d,
640 || self.label_texts.get(state.index).cloned().flatten(),
642 || {
644 self.values
645 .iter()
646 .enumerate()
647 .map(|(i, value_fn)| {
648 let name = self.names.get(i).cloned().unwrap_or_default();
649 Some((
650 self.series_stroke_from(&palette, i),
651 name,
652 value_fn(d).to_f64()?,
653 ))
654 })
655 .collect::<Option<Vec<_>>>()
656 },
657 window,
658 cx,
659 )?;
660
661 Some(tooltip.into_any_element())
662 }
663}
664
665#[cfg(test)]
666mod tests {
667 use super::*;
668
669 #[derive(Clone)]
670 struct Item {
671 subject: SharedString,
672 a: f64,
673 b: f64,
674 }
675
676 #[test]
677 fn test_radar_chart_builder() {
678 let data = vec![
679 Item {
680 subject: "Sales".into(),
681 a: 80.,
682 b: 60.,
683 },
684 Item {
685 subject: "Marketing".into(),
686 a: 50.,
687 b: 90.,
688 },
689 ];
690
691 let chart = RadarChart::new(data.clone())
692 .label(|d| d.subject.clone())
693 .value(|d| d.a)
694 .stroke(gpui::red())
695 .fill(gpui::red())
696 .name("A")
697 .value(|d| d.b)
698 .max_value(100.)
699 .outer_radius(120.)
700 .label_gap(8.)
701 .grid(false)
702 .grid_levels(5)
703 .dot()
704 .id("radar");
705
706 assert_eq!(chart.data.len(), 2);
707 assert_eq!(chart.values.len(), 2);
708 assert_eq!(chart.strokes.len(), 1);
709 assert_eq!(chart.fills.len(), 1);
710 assert_eq!(chart.names.len(), 1);
711 assert!(chart.label.is_some());
712 assert_eq!(chart.max_value, Some(100.));
713 assert_eq!(chart.outer_radius, 120.);
714 assert_eq!(chart.label_gap, 8.);
715 assert!(!chart.grid);
716 assert_eq!(chart.grid_levels, 5);
717 assert!(chart.dot);
718 assert_eq!(chart.id, gpui::ElementId::Name("radar".into()));
719
720 let values = (chart.values[0](&data[0]), chart.values[1](&data[0]));
721 assert_eq!(values, (80., 60.));
722 }
723
724 #[test]
727 fn test_radar_label_from_text() {
728 let labels = [
729 RadarLabel::from("Sales"),
730 RadarLabel::from("Sales".to_string()),
731 RadarLabel::from(SharedString::from("Sales")),
732 ];
733
734 for label in labels {
735 assert!(matches!(label, RadarLabel::Text(text) if text == "Sales"));
736 }
737 }
738
739 #[test]
740 fn test_radar_chart_grid_levels_min() {
741 let chart: RadarChart<Item, f64> = RadarChart::new(vec![]).grid_levels(0);
742 assert_eq!(chart.grid_levels, 1);
743 }
744
745 #[test]
746 fn test_radar_chart_hovered_index() {
747 let data = (0..4)
748 .map(|i| Item {
749 subject: format!("S{}", i).into(),
750 a: 50.,
751 b: 50.,
752 })
753 .collect::<Vec<_>>();
754
755 let chart: RadarChart<Item, f64> = RadarChart::new(data).value(|d| d.a);
758 let bounds = gpui::Bounds::new(point(px(0.), px(0.)), gpui::size(px(200.), px(200.)));
759
760 assert_eq!(
762 chart.hovered_index(point(px(100.), px(30.)), bounds),
763 Some(0)
764 );
765 assert_eq!(
766 chart.hovered_index(point(px(170.), px(100.)), bounds),
767 Some(1)
768 );
769 assert_eq!(
770 chart.hovered_index(point(px(100.), px(170.)), bounds),
771 Some(2)
772 );
773 assert_eq!(
774 chart.hovered_index(point(px(30.), px(100.)), bounds),
775 Some(3)
776 );
777
778 assert_eq!(
780 chart.hovered_index(point(px(110.), px(40.)), bounds),
781 Some(0)
782 );
783 assert_eq!(
784 chart.hovered_index(point(px(160.), px(90.)), bounds),
785 Some(1)
786 );
787
788 assert_eq!(chart.hovered_index(point(px(100.), px(5.)), bounds), None);
790 assert_eq!(chart.hovered_index(point(px(5.), px(5.)), bounds), None);
791 }
792}