Skip to main content

ggplot_rs/position/
stack.rs

1use std::borrow::Cow;
2use std::collections::HashMap;
3
4use crate::data::{DataFrame, Value};
5
6use super::{Position, PositionParams};
7
8/// Stack bars/areas on top of each other.
9pub struct PositionStack;
10
11impl Position for PositionStack {
12    fn compute(&self, data: &mut DataFrame, _params: &PositionParams) {
13        // Group by x, accumulate y values
14        let x_col = match data.column("x") {
15            Some(c) => c.to_vec(),
16            None => return,
17        };
18        let y_col = match data.column("y") {
19            Some(c) => c.to_vec(),
20            None => return,
21        };
22        super::preserve_raw_y(data, &y_col);
23
24        // ggplot2 stacks the first group at the TOP (so the stack order top-to-
25        // bottom matches the legend), so accumulate downward from each x's total
26        // rather than upward from 0.
27        // Per-x accumulators keyed by the borrowed x key: O(1) per row instead
28        // of a linear scan over the distinct x values.
29        let mut totals: HashMap<Cow<'_, str>, f64> = HashMap::new();
30        for (x, y) in x_col.iter().zip(y_col.iter()) {
31            *totals.entry(x.key_str()).or_insert(0.0) += y.as_f64().unwrap_or(0.0);
32        }
33
34        let mut consumed: HashMap<Cow<'_, str>, f64> = HashMap::new();
35        let mut new_y = Vec::with_capacity(y_col.len());
36        let mut ymin_vals = Vec::with_capacity(y_col.len());
37
38        for (x, y) in x_col.iter().zip(y_col.iter()) {
39            let x_key = x.key_str();
40            let y_val = y.as_f64().unwrap_or(0.0);
41            let total = totals.get(&x_key).copied().unwrap_or(0.0);
42            let run = consumed.get(&x_key).copied().unwrap_or(0.0);
43
44            // This group occupies [total - run - y, total - run] (top-down).
45            new_y.push(Value::Float(total - run));
46            ymin_vals.push(Value::Float(total - run - y_val));
47
48            *consumed.entry(x_key).or_insert(0.0) += y_val;
49        }
50
51        if let Some(col) = data.column_mut("y") {
52            *col = new_y;
53        }
54        if !data.has_column("ymin") {
55            data.add_column("ymin".to_string(), ymin_vals);
56        } else if let Some(col) = data.column_mut("ymin") {
57            *col = ymin_vals;
58        }
59    }
60
61    fn name(&self) -> &str {
62        "stack"
63    }
64}