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