Skip to main content

gpui_component/chart/
sankey_chart.rs

1use std::{
2    hash::{DefaultHasher, Hash, Hasher},
3    rc::Rc,
4};
5
6use gpui::{
7    AnyElement, App, Bounds, Corners, ElementId, Hsla, IntoElement, Pixels, Point, SharedString,
8    TextAlign, Window, fill, linear_color_stop, linear_gradient, point, prelude::FluentBuilder, px,
9};
10use gpui_component_macros::IntoPlot;
11
12use super::caller_id;
13use crate::{
14    ActiveTheme,
15    plot::{
16        PathCaches, Plot, ShapeKey,
17        label::{PlotLabel, TEXT_GAP, TEXT_SIZE, Text, measure_text_width, truncate_text_to_width},
18        origin_point,
19        shape::{
20            Sankey, SankeyAlign, SankeyGraph, SankeyLink, SankeyLinkLayout, SankeyValueScale,
21            sankey_link_path,
22        },
23        tooltip::{PlotHover, Tooltip, TooltipState},
24    },
25};
26
27const DEFAULT_NODE_WIDTH: f32 = 10.;
28const DEFAULT_NODE_PADDING: f32 = 16.;
29const DEFAULT_LINK_OPACITY: f32 = 0.3;
30const DEFAULT_MIN_LINK_WIDTH: f32 = 1.;
31const DEFAULT_LABEL_GAP: f32 = 6.;
32/// Cap each side's label margin (as a fraction of width) so a long label is
33/// truncated to a modest column beside the flow instead of dominating it.
34const MAX_LABEL_WIDTH_RATIO: f32 = 0.2;
35/// Cap the reserved top+bottom label band as a fraction of height.
36const MAX_LABEL_MARGIN_RATIO: f32 = 0.6;
37/// How much the links not attached to the hovered node fade, as a share of
38/// their opacity.
39const HOVER_DIM: f32 = 0.7;
40
41/// The placement of a sankey chart for one bounds size: the graph plus the
42/// label lines and margins it was laid out with.
43///
44/// Placing the graph relaxes the node order over several iterations, and the
45/// label margins need every label measured, so a chart keeps the frame in
46/// element state and reuses it while its key is unchanged.
47struct SankeyFrame {
48    graph: SankeyGraph,
49    layer_count: usize,
50    node_labels: Vec<Vec<SankeyLabel>>,
51    /// The label margins reserved on the left and right of the flow.
52    left: f32,
53    right: f32,
54}
55
56/// The frame of the last placement with the key it was placed for.
57#[derive(Default)]
58struct SankeyFrameCache {
59    key: Option<u64>,
60    frame: Option<Rc<SankeyFrame>>,
61}
62
63/// The hover a sankey chart paints, sampled once per frame in [`Plot::hover`].
64#[derive(Clone, Copy)]
65struct SankeyHover {
66    /// The hovered node.
67    node: usize,
68    /// How far the hover has faded in.
69    focus: f32,
70}
71
72/// A styled line of a sankey node label.
73#[derive(Clone)]
74pub struct SankeyLabel {
75    text: SharedString,
76    color: Option<Hsla>,
77    font_size: Option<f32>,
78}
79
80impl SankeyLabel {
81    /// Create a label line with the default color and font size.
82    pub fn new(text: impl Into<SharedString>) -> Self {
83        Self {
84            text: text.into(),
85            color: None,
86            font_size: None,
87        }
88    }
89
90    /// Set the text color. Defaults to the theme foreground.
91    pub fn color(mut self, color: impl Into<Hsla>) -> Self {
92        self.color = Some(color.into());
93        self
94    }
95
96    /// Set the font size. Defaults to 10.
97    pub fn font_size(mut self, font_size: f32) -> Self {
98        self.font_size = Some(font_size);
99        self
100    }
101
102    fn line_height(&self) -> f32 {
103        self.font_size.unwrap_or(TEXT_SIZE) + TEXT_GAP
104    }
105}
106
107fn block_height(lines: &[SankeyLabel]) -> f32 {
108    lines.iter().map(|line| line.line_height()).sum()
109}
110
111/// A Sankey diagram, layout modeled after [d3-sankey](https://github.com/d3/d3-sankey).
112///
113/// Links reference nodes by their index in the node list; map string ids to
114/// indices before constructing.
115#[derive(IntoPlot)]
116pub struct SankeyChart<T: 'static> {
117    nodes: Vec<T>,
118    links: Vec<SankeyLink>,
119    node_width: f32,
120    node_padding: f32,
121    align: SankeyAlign,
122    iterations: usize,
123    value_scale: SankeyValueScale,
124    node_corner_radius: Option<Pixels>,
125    node_color: Option<Rc<dyn Fn(&T) -> Hsla>>,
126    node_label: Option<Rc<dyn Fn(&T) -> SharedString>>,
127    value_label: Option<Rc<dyn Fn(&T, f64) -> SharedString>>,
128    labels: Option<Rc<dyn Fn(&T, f64) -> Vec<SankeyLabel>>>,
129    link_opacity: f32,
130    min_link_width: f32,
131    label_gap: f32,
132    tooltip_name: Option<Rc<dyn Fn(&T) -> SharedString + 'static>>,
133    tooltip_value: Option<Rc<dyn Fn(&T, f64) -> SharedString + 'static>>,
134    id: ElementId,
135    interactive: bool,
136    /// The placement for this frame, resolved in `prepaint` (measuring labels
137    /// needs the window) and read by `tooltip_state` and `paint`.
138    frame: Option<Rc<SankeyFrame>>,
139    hover: Option<SankeyHover>,
140}
141
142impl<T> SankeyChart<T> {
143    /// Create a chart from nodes and links; links reference nodes by their
144    /// index in `nodes` (map string ids to indices before constructing).
145    #[track_caller]
146    pub fn new<I, L>(nodes: I, links: L) -> Self
147    where
148        I: IntoIterator<Item = T>,
149        L: IntoIterator<Item = SankeyLink>,
150    {
151        Self {
152            nodes: nodes.into_iter().collect(),
153            links: links.into_iter().collect(),
154            node_width: DEFAULT_NODE_WIDTH,
155            node_padding: DEFAULT_NODE_PADDING,
156            align: SankeyAlign::default(),
157            iterations: 6,
158            value_scale: SankeyValueScale::default(),
159            node_corner_radius: None,
160            node_color: None,
161            node_label: None,
162            value_label: None,
163            labels: None,
164            link_opacity: DEFAULT_LINK_OPACITY,
165            min_link_width: DEFAULT_MIN_LINK_WIDTH,
166            label_gap: DEFAULT_LABEL_GAP,
167            tooltip_name: None,
168            tooltip_value: None,
169            id: caller_id(),
170            interactive: true,
171            frame: None,
172            hover: None,
173        }
174    }
175
176    /// Name this chart's [`ElementId`], replacing the default taken from the
177    /// construction site.
178    ///
179    /// Pass one where a single construction site renders several of these
180    /// charts as siblings: they share the default id, and with it one hover
181    /// state and one path cache. The id must be unique among those siblings.
182    pub fn id(mut self, id: impl Into<ElementId>) -> Self {
183        self.id = id.into();
184        self
185    }
186
187    /// Turn this chart's interactive layer on or off. On by default.
188    ///
189    /// The layer is the hitbox under the cursor and what it drives: the hovered
190    /// node's links stand out from the rest, and a tooltip shows its label and
191    /// throughput. Turn it off for a chart that only decorates, or one an element
192    /// above it wants the cursor for: without a hitbox it neither answers the
193    /// mouse nor takes the hover from what sits over it. A chart that is off also
194    /// drops its path cache, which is keyed on the same id.
195    pub fn interactive(mut self, interactive: bool) -> Self {
196        self.interactive = interactive;
197        self
198    }
199
200    /// Set the node rectangle width. Defaults to 10.
201    pub fn node_width(mut self, node_width: f32) -> Self {
202        self.node_width = node_width;
203        self
204    }
205
206    /// Set the vertical gap between nodes in a column. Defaults to 16.
207    pub fn node_padding(mut self, node_padding: f32) -> Self {
208        self.node_padding = node_padding;
209        self
210    }
211
212    /// Set the node alignment. Defaults to [`SankeyAlign::Justify`].
213    pub fn node_align(mut self, align: SankeyAlign) -> Self {
214        self.align = align;
215        self
216    }
217
218    /// Set the number of relaxation passes. Defaults to 6.
219    pub fn iterations(mut self, iterations: usize) -> Self {
220        self.iterations = iterations;
221        self
222    }
223
224    /// Set how flow values map to node heights.
225    ///
226    /// Defaults to [`SankeyValueScale::Linear`]. Use [`SankeyValueScale::Sqrt`]
227    /// to keep a dominant flow from dwarfing the small ones without
228    /// pre-transforming the data; labels still receive the raw values.
229    pub fn value_scale(mut self, value_scale: SankeyValueScale) -> Self {
230        self.value_scale = value_scale;
231        self
232    }
233
234    /// Set the corner radius of the node rectangles. Defaults to 0.
235    pub fn node_corner_radius(mut self, radius: impl Into<Pixels>) -> Self {
236        self.node_corner_radius = Some(radius.into());
237        self
238    }
239
240    /// Set the color of each node.
241    ///
242    /// Defaults to cycling the theme chart palette by node index.
243    pub fn node_color<H>(mut self, color: impl Fn(&T) -> H + 'static) -> Self
244    where
245        H: Into<Hsla> + 'static,
246    {
247        self.node_color = Some(Rc::new(move |t| color(t).into()));
248        self
249    }
250
251    /// Set the name label of each node, drawn in muted foreground. No name
252    /// label is drawn unless set.
253    pub fn node_label(mut self, label: impl Fn(&T) -> SharedString + 'static) -> Self {
254        self.node_label = Some(Rc::new(label));
255        self
256    }
257
258    /// Set the value label of each node, drawn above the name label. No value
259    /// label is drawn unless set.
260    ///
261    /// The closure receives the datum and the node's raw computed throughput
262    /// (max of incoming and outgoing flow, in unscaled units).
263    pub fn value_label(mut self, label: impl Fn(&T, f64) -> SharedString + 'static) -> Self {
264        self.value_label = Some(Rc::new(label));
265        self
266    }
267
268    /// Set fully custom node labels, one [`SankeyLabel`] per line, top to
269    /// bottom. Takes precedence over `node_label`/`value_label` when set;
270    /// unset by default.
271    ///
272    /// The closure receives the datum and the node's raw computed throughput
273    /// (max of incoming and outgoing flow, in unscaled units).
274    pub fn labels(mut self, labels: impl Fn(&T, f64) -> Vec<SankeyLabel> + 'static) -> Self {
275        self.labels = Some(Rc::new(labels));
276        self
277    }
278
279    /// Name the node under the cursor in the hover tooltip's row, beside its
280    /// throughput. Unset, the row carries no name at all.
281    ///
282    /// A sankey shows one number per node, so the node's own name is what the
283    /// row wants; without it the row reads as a swatch and a number with a gap
284    /// between them. The alternative was to title the tooltip from
285    /// `node_label`, but that also draws the name beside the node — and a chart
286    /// drawing its text through `labels` sets neither.
287    pub fn tooltip_name(mut self, name: impl Fn(&T) -> SharedString + 'static) -> Self {
288        self.tooltip_name = Some(Rc::new(name));
289        self
290    }
291
292    /// Set the text of the hover tooltip's row, the node's throughput.
293    ///
294    /// `value_label` supplies it when this is unset, and the raw number when
295    /// neither is set — which is what a chart drawing its text through `labels`
296    /// gets, however carefully it formats the value it draws.
297    pub fn tooltip_value(mut self, value: impl Fn(&T, f64) -> SharedString + 'static) -> Self {
298        self.tooltip_value = Some(Rc::new(value));
299        self
300    }
301
302    /// Set the opacity of the link ribbons. Defaults to 0.3.
303    pub fn link_opacity(mut self, opacity: f32) -> Self {
304        self.link_opacity = opacity;
305        self
306    }
307
308    /// Set the minimum ribbon thickness, so tiny flows stay visible. Defaults to 1.
309    pub fn min_link_width(mut self, width: f32) -> Self {
310        self.min_link_width = width;
311        self
312    }
313
314    /// Set the gap between a node and its labels. Defaults to 6.
315    pub fn label_gap(mut self, gap: f32) -> Self {
316        self.label_gap = gap;
317        self
318    }
319
320    fn sankey(&self) -> Sankey {
321        Sankey::new()
322            .node_width(self.node_width)
323            .node_padding(self.node_padding)
324            .node_align(self.align)
325            .iterations(self.iterations)
326            .value_scale(self.value_scale)
327    }
328
329    /// Raw per-node throughput (max of raw incoming and outgoing sums), for
330    /// labels — the layout's `node.value` is in scaled units under a
331    /// non-linear value scale, so labels must not use it.
332    fn raw_throughput(&self) -> Vec<f64> {
333        let mut incoming = vec![0f64; self.nodes.len()];
334        let mut outgoing = vec![0f64; self.nodes.len()];
335        for link in &self.links {
336            if let (Some(o), Some(i)) =
337                (outgoing.get_mut(link.source), incoming.get_mut(link.target))
338            {
339                *o += link.value;
340                *i += link.value;
341            }
342        }
343        incoming
344            .into_iter()
345            .zip(outgoing)
346            .map(|(i, o)| i.max(o))
347            .collect()
348    }
349}
350
351impl<T> SankeyChart<T> {
352    /// Each node's label lines: the custom `labels` closure wins, otherwise the
353    /// value/name lines with the default styles. Labels get the raw throughput,
354    /// not the layout's (possibly scaled) value.
355    fn node_labels(&self, cx: &App) -> Vec<Vec<SankeyLabel>> {
356        let raw_value = self.raw_throughput();
357        self.nodes
358            .iter()
359            .zip(raw_value)
360            .map(|(datum, value)| {
361                if let Some(labels) = &self.labels {
362                    labels(datum, value)
363                } else {
364                    let mut lines = Vec::new();
365                    if let Some(value_label) = &self.value_label {
366                        lines.push(SankeyLabel::new(value_label(datum, value)));
367                    }
368                    if let Some(node_label) = &self.node_label {
369                        lines.push(
370                            SankeyLabel::new(node_label(datum)).color(cx.theme().muted_foreground),
371                        );
372                    }
373                    lines
374                }
375            })
376            .collect()
377    }
378
379    /// The key a placement is reused under: everything that shapes it, which is
380    /// the bounds size, the graph, the placement settings and the label lines.
381    fn frame_key(&self, bounds: Bounds<Pixels>, node_labels: &[Vec<SankeyLabel>]) -> u64 {
382        let mut hasher = DefaultHasher::new();
383        bounds.size.width.as_f32().to_bits().hash(&mut hasher);
384        bounds.size.height.as_f32().to_bits().hash(&mut hasher);
385        self.nodes.len().hash(&mut hasher);
386        for link in &self.links {
387            link.source.hash(&mut hasher);
388            link.target.hash(&mut hasher);
389            link.value.to_bits().hash(&mut hasher);
390        }
391        self.node_width.to_bits().hash(&mut hasher);
392        self.node_padding.to_bits().hash(&mut hasher);
393        self.align.hash(&mut hasher);
394        self.iterations.hash(&mut hasher);
395        self.value_scale.hash(&mut hasher);
396        self.label_gap.to_bits().hash(&mut hasher);
397        for lines in node_labels {
398            lines.len().hash(&mut hasher);
399            for line in lines {
400                line.text.hash(&mut hasher);
401                line.font_size.map(f32::to_bits).hash(&mut hasher);
402                line.color
403                    .map(|color| [color.h, color.s, color.l, color.a].map(f32::to_bits))
404                    .hash(&mut hasher);
405            }
406        }
407        hasher.finish()
408    }
409
410    /// Place the graph within `bounds`, reserving margins for the labels.
411    fn place(
412        &self,
413        bounds: Bounds<Pixels>,
414        node_labels: Vec<Vec<SankeyLabel>>,
415        window: &mut Window,
416    ) -> Option<SankeyFrame> {
417        let width = bounds.size.width.as_f32();
418        let height = bounds.size.height.as_f32();
419
420        // First pass: only the topology (each node's `layer`) is needed to
421        // measure the label margins.
422        let topology = self.sankey().topology(self.nodes.len(), &self.links).ok()?;
423        let layer_count = topology.layer_count();
424        let has_labels = node_labels.iter().any(|lines| !lines.is_empty());
425
426        // Reserve margins so the labels beside the first/last columns and
427        // above the middle columns are not clipped.
428        let mut left = 0f32;
429        let mut right = 0f32;
430        if has_labels {
431            for node in &topology.nodes {
432                if node.layer != 0 && node.layer + 1 != layer_count {
433                    continue;
434                }
435                let mut label_width = 0f32;
436                for line in &node_labels[node.index] {
437                    label_width = label_width.max(measure_text_width(
438                        &line.text,
439                        px(line.font_size.unwrap_or(TEXT_SIZE)),
440                        window,
441                    ));
442                }
443                if node.layer == 0 {
444                    left = left.max(label_width + self.label_gap);
445                } else {
446                    right = right.max(label_width + self.label_gap);
447                }
448            }
449
450            // Cap each side independently so one long label is truncated to a
451            // modest column rather than eating into the flow area.
452            let side_cap = width * MAX_LABEL_WIDTH_RATIO;
453            left = left.min(side_cap);
454            right = right.min(side_cap);
455        }
456        // Above-node labels are only emitted for middle columns, so reserve
457        // the top band for the tallest such label block. Cap the vertical
458        // margins like the horizontal ones so a short chart doesn't collapse
459        // the flow.
460        let mut top = 0f32;
461        if has_labels && layer_count > 2 {
462            for node in &topology.nodes {
463                if node.layer == 0 || node.layer + 1 == layer_count {
464                    continue;
465                }
466                let block = block_height(&node_labels[node.index]);
467                if block > 0. {
468                    top = top.max(block + TEXT_GAP);
469                }
470            }
471        }
472        let mut bottom = if has_labels { TEXT_GAP } else { 0. };
473        let max_vertical = height * MAX_LABEL_MARGIN_RATIO;
474        if top + bottom > max_vertical {
475            let k = max_vertical / (top + bottom);
476            top *= k;
477            bottom *= k;
478        }
479
480        // Second pass: complete the placement on the final extent, reusing
481        // the first pass's topology.
482        let graph = self
483            .sankey()
484            .extent(
485                left,
486                top,
487                (width - right).max(left + 1.),
488                (height - bottom).max(top + 1.),
489            )
490            .layout_from(topology);
491
492        Some(SankeyFrame {
493            graph,
494            layer_count,
495            node_labels,
496            left,
497            right,
498        })
499    }
500
501    /// Whether `link` starts or ends at `node`.
502    fn is_attached(link: &SankeyLinkLayout, node: usize) -> bool {
503        link.source == node || link.target == node
504    }
505}
506
507impl<T> Plot for SankeyChart<T> {
508    /// Resolve the placement for the frame, reusing the last one while nothing
509    /// that shapes it has changed. Measuring the labels needs the window, which
510    /// `tooltip_state` does not have, so this runs here rather than in `paint`.
511    fn prepaint(
512        &mut self,
513        bounds: Bounds<Pixels>,
514        window: &mut Window,
515        cx: &mut App,
516    ) -> Vec<AnyElement> {
517        self.frame = None;
518        let width = bounds.size.width.as_f32();
519        let height = bounds.size.height.as_f32();
520        if self.nodes.is_empty() || self.links.is_empty() || width <= 0. || height <= 0. {
521            return vec![];
522        }
523
524        let node_labels = self.node_labels(cx);
525
526        // Caching hangs off the chart's own id, which only an interactive chart
527        // puts on the stack; without one, siblings would share a slot and thrash
528        // it, so a chart that is off places itself afresh each paint.
529        self.frame = if self.interactive {
530            let key = self.frame_key(bounds, &node_labels);
531            let cache =
532                window.use_keyed_state("sankey-frame", cx, |_, _| SankeyFrameCache::default());
533            let cached = cache.read(cx);
534            if cached.key == Some(key) {
535                cached.frame.clone()
536            } else {
537                let frame = self.place(bounds, node_labels, window).map(Rc::new);
538                cache.update(cx, |cache, _| {
539                    cache.key = Some(key);
540                    cache.frame = frame.clone();
541                });
542                frame
543            }
544        } else {
545            self.place(bounds, node_labels, window).map(Rc::new)
546        };
547
548        vec![]
549    }
550
551    fn paint(&mut self, bounds: Bounds<Pixels>, window: &mut Window, cx: &mut App) {
552        let Some(frame) = self.frame.clone() else {
553            return;
554        };
555        let SankeyFrame {
556            graph,
557            layer_count,
558            node_labels,
559            left,
560            right,
561        } = &*frame;
562        let (layer_count, left, right) = (*layer_count, *left, *right);
563        let width = bounds.size.width.as_f32();
564        let height = bounds.size.height.as_f32();
565
566        let palette = [
567            cx.theme().chart_1,
568            cx.theme().chart_2,
569            cx.theme().chart_3,
570            cx.theme().chart_4,
571            cx.theme().chart_5,
572        ];
573        let colors: Vec<Hsla> = self
574            .nodes
575            .iter()
576            .enumerate()
577            .map(|(index, datum)| match &self.node_color {
578                Some(color) => color(datum),
579                None => palette[index % palette.len()],
580            })
581            .collect();
582
583        // Links first, under the nodes. The links of the hovered node keep their
584        // opacity while the rest fade behind them.
585        //
586        // Hovering changes only a ribbon's opacity, so an interactive chart keeps
587        // each tessellated ribbon, slotted by the link's index in the graph so a
588        // skipped zero-value link doesn't shift the others. Without an id, siblings
589        // would share the slots, so a chart that is off tessellates afresh.
590        let min_width = self.min_link_width;
591        let caches = self
592            .interactive
593            .then(|| PathCaches::for_paint("links", window, cx));
594        for (ix, link) in graph.links.iter().enumerate() {
595            if link.value <= 0. {
596                continue;
597            }
598            let source = &graph.nodes[link.source];
599            let target = &graph.nodes[link.target];
600            let path = match caches.as_ref() {
601                Some(caches) => caches.update(cx, |caches, _| {
602                    let key = ShapeKey::new(())
603                        .f32(source.x1)
604                        .f32(target.x0)
605                        .f32(link.y0)
606                        .f32(link.y1)
607                        .f32(link.source_width.max(min_width))
608                        .f32(link.target_width.max(min_width))
609                        .finish();
610                    caches.slot(ix).get(key, bounds.origin, || {
611                        sankey_link_path(source, target, link, min_width, Point::default())
612                    })
613                }),
614                None => sankey_link_path(source, target, link, min_width, bounds.origin),
615            };
616            let Some(path) = path else {
617                continue;
618            };
619            let opacity = match self.hover {
620                Some(hover) if !Self::is_attached(link, hover.node) => {
621                    self.link_opacity * (1. - HOVER_DIM * hover.focus)
622                }
623                _ => self.link_opacity,
624            };
625            window.paint_path(
626                path,
627                linear_gradient(
628                    90.,
629                    linear_color_stop(colors[link.source].opacity(opacity), 0.),
630                    linear_color_stop(colors[link.target].opacity(opacity), 1.),
631                ),
632            );
633        }
634
635        let corner_radii = Corners::all(self.node_corner_radius.unwrap_or_default());
636        for node in &graph.nodes {
637            let node_bounds = Bounds::from_corners(
638                origin_point(px(node.x0), px(node.y0), bounds.origin),
639                // Keep tiny nodes visible with a minimum 1px height.
640                origin_point(px(node.x1), px(node.y1.max(node.y0 + 1.)), bounds.origin),
641            );
642            window.paint_quad(fill(node_bounds, colors[node.index]).corner_radii(corner_radii));
643        }
644
645        let mut texts = Vec::new();
646        for node in &graph.nodes {
647            let lines = &node_labels[node.index];
648            if lines.is_empty() {
649                continue;
650            }
651
652            let is_first = node.layer == 0;
653            let is_last = node.layer + 1 == layer_count;
654            // `x`/`align` place the label beside (first/last) or centered above
655            // (middle) the node, and `max_width` bounds it so a long label is
656            // truncated with an ellipsis instead of drawn outside the plot:
657            // first/last to their reserved margin, middle to twice the smaller
658            // gap to the plot edge (generous for interior nodes, only bites a
659            // label long enough to actually run off-plot).
660            let (x, align, max_width) = if is_first {
661                (
662                    node.x0 - self.label_gap,
663                    TextAlign::Right,
664                    left - self.label_gap,
665                )
666            } else if is_last {
667                (
668                    node.x1 + self.label_gap,
669                    TextAlign::Left,
670                    right - self.label_gap,
671                )
672            } else {
673                let center = (node.x0 + node.x1) / 2.;
674                let edge_budget = 2. * center.min(width - center);
675                (center, TextAlign::Center, edge_budget)
676            };
677
678            let block = block_height(lines);
679            let mut y = if is_first || is_last {
680                // Block vertically centered beside the node, clamped into
681                // the plot area so labels of nodes near the top or bottom
682                // edge are not clipped.
683                ((node.y0 + node.y1) / 2. - block / 2.)
684                    .min(height - block)
685                    .max(0.)
686            } else {
687                // Block above the node.
688                node.y0 - block - TEXT_GAP
689            };
690
691            for line in lines {
692                let font_size = px(line.font_size.unwrap_or(TEXT_SIZE));
693                let text = truncate_text_to_width(&line.text, font_size, max_width, window);
694                texts.push(
695                    Text::new(
696                        text,
697                        point(px(x), px(y)),
698                        line.color.unwrap_or(cx.theme().foreground),
699                    )
700                    .font_size(font_size)
701                    .align(align),
702                );
703                y += line.line_height();
704            }
705        }
706        PlotLabel::new(texts).paint(&bounds, window, cx);
707    }
708
709    fn id(&self) -> Option<ElementId> {
710        self.interactive.then(|| self.id.clone())
711    }
712
713    fn tooltip_state(
714        &self,
715        position: Point<Pixels>,
716        _bounds: Bounds<Pixels>,
717        _cx: &App,
718    ) -> Option<TooltipState> {
719        let frame = self.frame.as_ref()?;
720        let (x, y) = (position.x.as_f32(), position.y.as_f32());
721        let node = frame.graph.nodes.iter().find(|node| {
722            (node.x0..=node.x1).contains(&x) && (node.y0..=node.y1.max(node.y0 + 1.)).contains(&y)
723        })?;
724        Some(TooltipState::new(node.index, position, vec![]))
725    }
726
727    fn hover(&mut self, hover: Option<&PlotHover>, _window: &mut Window, _cx: &mut App) {
728        self.hover = hover.map(|hover| SankeyHover {
729            node: hover.state().index,
730            focus: hover.progress(),
731        });
732    }
733
734    fn tooltip(
735        &self,
736        state: &TooltipState,
737        cursor: Point<Pixels>,
738        bounds: Bounds<Pixels>,
739        _window: &mut Window,
740        cx: &mut App,
741    ) -> Option<AnyElement> {
742        let datum = self.nodes.get(state.index)?;
743        let value = self.raw_throughput().get(state.index).copied()?;
744        let color = match &self.node_color {
745            Some(color) => color(datum),
746            None => {
747                let palette = [
748                    cx.theme().chart_1,
749                    cx.theme().chart_2,
750                    cx.theme().chart_3,
751                    cx.theme().chart_4,
752                    cx.theme().chart_5,
753                ];
754                palette[state.index % palette.len()]
755            }
756        };
757        let value_text = match self.tooltip_value.as_ref().or(self.value_label.as_ref()) {
758            Some(value_text) => value_text(datum, value),
759            None => format!("{value}").into(),
760        };
761
762        Some(
763            // Follow the cursor; the node's links mark it.
764            Tooltip::new(cursor, bounds.size)
765                .gap(px(8.))
766                .when_some(self.node_label.as_ref(), |this, label| {
767                    this.title(label(datum))
768                })
769                .row(
770                    color,
771                    match self.tooltip_name.as_ref() {
772                        Some(tooltip_name) => tooltip_name(datum),
773                        None => SharedString::default(),
774                    },
775                    value_text,
776                )
777                .into_any_element(),
778        )
779    }
780}
781
782#[cfg(test)]
783mod tests {
784    use super::*;
785
786    #[test]
787    fn test_sankey_chart_builder() {
788        let chart = SankeyChart::new(vec!["a", "b"], vec![SankeyLink::new(0, 1, 5.)]);
789        assert_eq!(chart.nodes.len(), 2);
790        assert_eq!(chart.links.len(), 1);
791        assert_eq!(chart.node_width, DEFAULT_NODE_WIDTH);
792        assert_eq!(chart.node_padding, DEFAULT_NODE_PADDING);
793        assert_eq!(chart.align, SankeyAlign::Justify);
794        assert_eq!(chart.iterations, 6);
795        assert_eq!(chart.node_corner_radius, None);
796        assert_eq!(chart.link_opacity, DEFAULT_LINK_OPACITY);
797        assert_eq!(chart.min_link_width, DEFAULT_MIN_LINK_WIDTH);
798        assert_eq!(chart.label_gap, DEFAULT_LABEL_GAP);
799        assert!(chart.node_color.is_none());
800        assert!(chart.node_label.is_none());
801        assert!(chart.value_label.is_none());
802        assert!(chart.labels.is_none());
803
804        let chart = chart
805            .node_width(8.)
806            .node_padding(20.)
807            .node_align(SankeyAlign::Left)
808            .iterations(10)
809            .node_corner_radius(px(2.))
810            .node_color(|_| gpui::red())
811            .node_label(|d| SharedString::from(d.to_string()))
812            .value_label(|_, value| SharedString::from(format!("{}", value)))
813            .labels(|d, value| {
814                vec![
815                    SankeyLabel::new(format!("{}", value)),
816                    SankeyLabel::new(d.to_string()),
817                ]
818            })
819            .link_opacity(0.5)
820            .min_link_width(2.)
821            .label_gap(10.);
822        assert_eq!(chart.node_width, 8.);
823        assert_eq!(chart.node_padding, 20.);
824        assert_eq!(chart.align, SankeyAlign::Left);
825        assert_eq!(chart.iterations, 10);
826        assert_eq!(chart.node_corner_radius, Some(px(2.)));
827        assert_eq!(chart.link_opacity, 0.5);
828        assert_eq!(chart.min_link_width, 2.);
829        assert_eq!(chart.label_gap, 10.);
830        assert!(chart.node_color.is_some());
831        assert!(chart.node_label.is_some());
832        assert!(chart.value_label.is_some());
833        assert!(chart.labels.is_some());
834    }
835
836    #[test]
837    fn test_sankey_label_builder() {
838        let label = SankeyLabel::new("a");
839        assert_eq!(label.text, "a");
840        assert_eq!(label.color, None);
841        assert_eq!(label.font_size, None);
842        assert_eq!(label.line_height(), TEXT_SIZE + TEXT_GAP);
843
844        let label = SankeyLabel::new("b").color(gpui::red()).font_size(14.);
845        assert_eq!(label.color, Some(gpui::red()));
846        assert_eq!(label.font_size, Some(14.));
847        assert_eq!(label.line_height(), 14. + TEXT_GAP);
848
849        assert_eq!(
850            block_height(&[SankeyLabel::new("a"), SankeyLabel::new("b").font_size(14.)]),
851            TEXT_SIZE + TEXT_GAP + 14. + TEXT_GAP
852        );
853        assert_eq!(block_height(&[]), 0.);
854    }
855
856    #[test]
857    fn test_sankey_chart_raw_throughput() {
858        // A(out 30) -> B, B -> C(20) + D(10): B's throughput is max(in, out).
859        let chart = SankeyChart::new(
860            vec!["a", "b", "c", "d"],
861            vec![
862                SankeyLink::new(0, 1, 30.),
863                SankeyLink::new(1, 2, 20.),
864                SankeyLink::new(1, 3, 10.),
865            ],
866        );
867        let raw = chart.raw_throughput();
868        assert_eq!(raw, vec![30., 30., 20., 10.]);
869
870        // Under Sqrt the layout's node value is scaled, but raw_throughput
871        // (used for labels) must stay in raw units — the two must differ.
872        let sqrt = chart
873            .value_scale(SankeyValueScale::Sqrt)
874            .sankey()
875            .layout(4, &chart_links())
876            .unwrap();
877        // Node A: layout value is sqrt-scaled (30 -> sqrt(30)), raw is 30.
878        assert!((sqrt.nodes[0].value - 30f64.sqrt()).abs() < 1e-6);
879        assert!((raw[0] - 30.).abs() < 1e-6);
880        assert!(raw[0] != sqrt.nodes[0].value);
881    }
882
883    fn chart_links() -> Vec<SankeyLink> {
884        vec![
885            SankeyLink::new(0, 1, 30.),
886            SankeyLink::new(1, 2, 20.),
887            SankeyLink::new(1, 3, 10.),
888        ]
889    }
890
891    /// A chart drawing its text through `labels` sets neither `node_label` nor
892    /// `value_label`, so its tooltip row had no name and an unformatted number.
893    #[test]
894    fn test_tooltip_text_is_settable_without_drawing_labels() {
895        let chart = SankeyChart::new(vec!["Revenue"], Vec::<SankeyLink>::new())
896            .tooltip_name(|_| "Revenue".into())
897            .tooltip_value(|_, value| format!("{value:.0}M").into());
898        assert!(chart.tooltip_name.is_some());
899        assert!(chart.tooltip_value.is_some());
900        assert!(chart.node_label.is_none());
901        assert!(chart.value_label.is_none());
902    }
903}