Skip to main content

ggplot_rs/data/
dataframe.rs

1use indexmap::IndexMap;
2
3use super::{GroupKey, Value};
4
5/// Internal columnar DataFrame for data storage and manipulation.
6#[derive(Clone, Debug)]
7pub struct DataFrame {
8    columns: IndexMap<String, Vec<Value>>,
9    nrows: usize,
10    /// Problems found while assembling the frame (e.g. mismatched column
11    /// lengths). Reported by [`validate`](Self::validate), and turned into a
12    /// `GGError::ValidationError` when a plot using this frame is built.
13    issues: Vec<String>,
14}
15
16impl DataFrame {
17    /// Create an empty DataFrame.
18    pub fn new() -> Self {
19        DataFrame {
20            columns: IndexMap::new(),
21            nrows: 0,
22            issues: Vec::new(),
23        }
24    }
25
26    /// Get a column by name.
27    pub fn column(&self, name: &str) -> Option<&[Value]> {
28        self.columns.get(name).map(|v| v.as_slice())
29    }
30
31    /// Get number of rows.
32    pub fn nrows(&self) -> usize {
33        self.nrows
34    }
35
36    /// Get number of columns.
37    pub fn ncols(&self) -> usize {
38        self.columns.len()
39    }
40
41    /// Get column names.
42    pub fn column_names(&self) -> Vec<&str> {
43        self.columns.keys().map(|s| s.as_str()).collect()
44    }
45
46    /// Check if a column exists.
47    pub fn has_column(&self, name: &str) -> bool {
48        self.columns.contains_key(name)
49    }
50
51    /// Add (or replace) a column.
52    ///
53    /// Never panics. If the length doesn't match the existing rows, the shorter
54    /// side is padded with [`Value::Na`] and the mismatch is recorded as an
55    /// [issue](Self::issues): building or rendering a plot from this frame then
56    /// fails with `GGError::ValidationError`. Use
57    /// [`try_add_column`](Self::try_add_column) to reject the column instead.
58    pub fn add_column(&mut self, name: String, mut values: Vec<Value>) {
59        if self.columns.is_empty() {
60            self.nrows = values.len();
61        } else if values.len() != self.nrows {
62            self.issues
63                .push(Self::mismatch_message(&name, values.len(), self.nrows));
64            if values.len() < self.nrows {
65                values.resize(self.nrows, Value::Na);
66            } else {
67                let n = values.len();
68                for col in self.columns.values_mut() {
69                    col.resize(n, Value::Na);
70                }
71                self.nrows = n;
72            }
73        }
74        self.columns.insert(name, values);
75    }
76
77    /// Add (or replace) a column, rejecting a length mismatch with a
78    /// `GGError::ValidationError` (the frame is left unchanged).
79    pub fn try_add_column(
80        &mut self,
81        name: String,
82        values: Vec<Value>,
83    ) -> Result<(), crate::plot::GGError> {
84        if !self.columns.is_empty() && values.len() != self.nrows {
85            return Err(crate::plot::GGError::ValidationError(
86                Self::mismatch_message(&name, values.len(), self.nrows),
87            ));
88        }
89        self.add_column(name, values);
90        Ok(())
91    }
92
93    fn mismatch_message(name: &str, len: usize, nrows: usize) -> String {
94        format!("column '{name}' has {len} values but the data has {nrows} rows")
95    }
96
97    /// Problems recorded while assembling this frame (empty when well-formed).
98    pub fn issues(&self) -> &[String] {
99        &self.issues
100    }
101
102    /// `Ok` when the frame is well-formed, else a `GGError::ValidationError`
103    /// describing every recorded [issue](Self::issues).
104    pub fn validate(&self) -> Result<(), crate::plot::GGError> {
105        if self.issues.is_empty() {
106            Ok(())
107        } else {
108            Err(crate::plot::GGError::ValidationError(
109                self.issues.join("; "),
110            ))
111        }
112    }
113
114    /// Get a mutable reference to a column.
115    pub fn column_mut(&mut self, name: &str) -> Option<&mut Vec<Value>> {
116        self.columns.get_mut(name)
117    }
118
119    /// Group by one or more key columns. Returns a Vec of DataFrames, one per group.
120    pub fn group_by(&self, keys: &[&str]) -> Vec<DataFrame> {
121        if self.nrows == 0 {
122            return vec![];
123        }
124
125        // Build group keys for each row
126        // Borrowing keys (no per-row String clones for text columns); a missing
127        // value is its own group, distinct from the literal string "NA".
128        let key_cols: Vec<Option<&Vec<Value>>> =
129            keys.iter().map(|k| self.columns.get(*k)).collect();
130        let mut group_map: IndexMap<Vec<GroupKey<'_>>, Vec<usize>> = IndexMap::new();
131
132        for i in 0..self.nrows {
133            let key: Vec<GroupKey<'_>> = key_cols
134                .iter()
135                .map(|col| col.map_or(GroupKey::Na, |c| c[i].group_key()))
136                .collect();
137            group_map.entry(key).or_default().push(i);
138        }
139
140        group_map
141            .into_values()
142            .map(|indices| {
143                let mut df = DataFrame::new();
144                for (name, col) in &self.columns {
145                    let values: Vec<Value> = indices.iter().map(|&i| col[i].clone()).collect();
146                    df.add_column(name.clone(), values);
147                }
148                df
149            })
150            .collect()
151    }
152
153    /// Vertically stack another DataFrame onto this one.
154    pub fn vstack(&mut self, other: &DataFrame) {
155        if other.nrows == 0 {
156            return;
157        }
158        if self.columns.is_empty() {
159            *self = other.clone();
160            return;
161        }
162
163        // Add columns from other that we have
164        for (name, col) in &self.columns {
165            if let Some(other_col) = other.columns.get(name) {
166                // Will extend below
167                let _ = (col, other_col);
168            }
169        }
170
171        // Also add columns from other that we don't have (fill with NA)
172        for name in other.columns.keys() {
173            if !self.columns.contains_key(name) {
174                self.columns
175                    .insert(name.clone(), vec![Value::Na; self.nrows]);
176            }
177        }
178
179        let old_nrows = self.nrows;
180        self.nrows += other.nrows;
181
182        for (name, col) in &mut self.columns {
183            if let Some(other_col) = other.columns.get(name) {
184                col.extend(other_col.iter().cloned());
185            } else {
186                col.extend(std::iter::repeat_with(|| Value::Na).take(other.nrows));
187            }
188            debug_assert_eq!(col.len(), old_nrows + other.nrows);
189        }
190    }
191
192    /// Select a subset of columns.
193    pub fn select(&self, columns: &[&str]) -> DataFrame {
194        let mut df = DataFrame::new();
195        for &col_name in columns {
196            if let Some(col) = self.columns.get(col_name) {
197                df.add_column(col_name.to_string(), col.clone());
198            }
199        }
200        df
201    }
202
203    /// Get a single row as a map.
204    pub fn row(&self, idx: usize) -> IndexMap<String, Value> {
205        assert!(
206            idx < self.nrows,
207            "Row index {idx} out of bounds ({} rows)",
208            self.nrows
209        );
210        let mut map = IndexMap::new();
211        for (name, col) in &self.columns {
212            map.insert(name.clone(), col[idx].clone());
213        }
214        map
215    }
216
217    /// Sort by a column (ascending). Returns a new DataFrame.
218    pub fn sort_by(&self, column: &str) -> DataFrame {
219        let col = match self.columns.get(column) {
220            Some(c) => c,
221            None => return self.clone(),
222        };
223
224        let mut indices: Vec<usize> = (0..self.nrows).collect();
225        indices.sort_by(|&a, &b| {
226            let va = col[a].as_f64().unwrap_or(f64::NAN);
227            let vb = col[b].as_f64().unwrap_or(f64::NAN);
228            va.total_cmp(&vb)
229        });
230
231        let mut df = DataFrame::new();
232        for (name, c) in &self.columns {
233            let values: Vec<Value> = indices.iter().map(|&i| c[i].clone()).collect();
234            df.add_column(name.clone(), values);
235        }
236        df
237    }
238
239    /// Create from rows (list of maps).
240    pub fn from_rows(rows: Vec<IndexMap<String, Value>>) -> Self {
241        if rows.is_empty() {
242            return DataFrame::new();
243        }
244
245        // Collect all column names from all rows
246        let mut col_names: IndexMap<String, ()> = IndexMap::new();
247        for row in &rows {
248            for key in row.keys() {
249                col_names.entry(key.clone()).or_default();
250            }
251        }
252
253        let mut df = DataFrame::new();
254        for name in col_names.keys() {
255            let values: Vec<Value> = rows
256                .iter()
257                .map(|row| row.get(name).cloned().unwrap_or(Value::Na))
258                .collect();
259            df.add_column(name.clone(), values);
260        }
261        df
262    }
263
264    /// Get all unique values in a column.
265    pub fn unique_values(&self, column: &str) -> Vec<Value> {
266        let col = match self.columns.get(column) {
267            Some(c) => c,
268            None => return vec![],
269        };
270        let mut seen = std::collections::HashSet::new();
271        let mut result = Vec::new();
272        for v in col {
273            if seen.insert(v.group_key()) {
274                result.push(v.clone());
275            }
276        }
277        result
278    }
279}
280
281impl DataFrame {
282    /// Load a DataFrame from a CSV file.
283    /// First row is treated as column headers.
284    /// Values are parsed as Float if possible, otherwise kept as strings.
285    /// The literal string "NA" is parsed as Value::Na.
286    pub fn from_csv(path: &str) -> Result<Self, std::io::Error> {
287        let content = std::fs::read_to_string(path)?;
288        let mut lines = content.lines();
289
290        let header = match lines.next() {
291            Some(h) => h,
292            None => return Ok(DataFrame::new()),
293        };
294
295        let col_names: Vec<&str> = header.split(',').map(|s| s.trim()).collect();
296        let mut columns: Vec<Vec<Value>> = vec![Vec::new(); col_names.len()];
297
298        for line in lines {
299            let line = line.trim();
300            if line.is_empty() {
301                continue;
302            }
303            let fields: Vec<&str> = line.split(',').collect();
304            for (i, field) in fields.iter().enumerate() {
305                if i >= col_names.len() {
306                    continue;
307                }
308                let field = field.trim();
309                let val = if field == "NA" || field == "na" {
310                    Value::Na
311                } else if let Ok(f) = field.parse::<f64>() {
312                    Value::Float(f)
313                } else {
314                    Value::Str(field.to_string())
315                };
316                columns[i].push(val);
317            }
318            // Pad missing columns with NA
319            for col in columns.iter_mut().skip(fields.len()) {
320                col.push(Value::Na);
321            }
322        }
323
324        let mut df = DataFrame::new();
325        for (i, name) in col_names.iter().enumerate() {
326            if !columns[i].is_empty() {
327                df.add_column(name.to_string(), std::mem::take(&mut columns[i]));
328            }
329        }
330
331        Ok(df)
332    }
333}
334
335impl Default for DataFrame {
336    fn default() -> Self {
337        Self::new()
338    }
339}
340
341#[cfg(test)]
342mod tests {
343    use super::*;
344
345    #[test]
346    fn mismatched_column_lengths_pad_and_record_an_issue() {
347        let mut df = DataFrame::new();
348        df.add_column("x".into(), vec![Value::Float(1.0), Value::Float(2.0)]);
349        df.add_column("y".into(), vec![Value::Float(3.0)]);
350        assert_eq!(df.nrows(), 2);
351        assert_eq!(df.column("y").unwrap()[1], Value::Na);
352        df.add_column("z".into(), vec![Value::Float(0.0); 4]);
353        assert_eq!(df.nrows(), 4);
354        assert!(df
355            .column_names()
356            .iter()
357            .all(|c| df.column(c).unwrap().len() == 4));
358        assert_eq!(df.issues().len(), 2);
359        let err = df.validate().unwrap_err().to_string();
360        assert!(err.contains("'y'") && err.contains("'z'"), "{err}");
361
362        let mut ok = DataFrame::new();
363        ok.add_column("x".into(), vec![Value::Float(1.0)]);
364        assert!(ok.validate().is_ok());
365        assert!(ok
366            .try_add_column("y".into(), vec![Value::Na, Value::Na])
367            .is_err());
368        assert!(!ok.has_column("y"), "rejected column is not added");
369        assert!(ok.try_add_column("y".into(), vec![Value::Na]).is_ok());
370        assert!(ok.validate().is_ok());
371    }
372
373    #[test]
374    fn group_by_keeps_na_apart_from_literal_na() {
375        let mut df = DataFrame::new();
376        df.add_column(
377            "g".into(),
378            vec![
379                Value::Na,
380                Value::Str("NA".into()),
381                Value::Na,
382                Value::Str("NA".into()),
383            ],
384        );
385        df.add_column("v".into(), (0..4).map(Value::Integer).collect());
386        let groups = df.group_by(&["g"]);
387        assert_eq!(groups.len(), 2);
388        assert!(groups[0].column("g").unwrap().iter().all(Value::is_na));
389        assert!(groups[1]
390            .column("g")
391            .unwrap()
392            .iter()
393            .all(|v| v.as_str() == Some("NA")));
394        assert_eq!(df.unique_values("g").len(), 2);
395    }
396
397    #[test]
398    fn test_add_column_and_access() {
399        let mut df = DataFrame::new();
400        df.add_column("x".into(), vec![Value::Float(1.0), Value::Float(2.0)]);
401        df.add_column("y".into(), vec![Value::Float(3.0), Value::Float(4.0)]);
402
403        assert_eq!(df.nrows(), 2);
404        assert_eq!(df.ncols(), 2);
405        assert!(df.has_column("x"));
406        assert!(!df.has_column("z"));
407    }
408
409    #[test]
410    fn test_group_by() {
411        let mut df = DataFrame::new();
412        df.add_column(
413            "cat".into(),
414            vec![
415                Value::Str("a".into()),
416                Value::Str("b".into()),
417                Value::Str("a".into()),
418            ],
419        );
420        df.add_column(
421            "val".into(),
422            vec![Value::Float(1.0), Value::Float(2.0), Value::Float(3.0)],
423        );
424
425        let groups = df.group_by(&["cat"]);
426        assert_eq!(groups.len(), 2);
427        assert_eq!(groups[0].nrows(), 2); // "a" group
428        assert_eq!(groups[1].nrows(), 1); // "b" group
429    }
430
431    #[test]
432    fn test_vstack() {
433        let mut df1 = DataFrame::new();
434        df1.add_column("x".into(), vec![Value::Float(1.0)]);
435
436        let mut df2 = DataFrame::new();
437        df2.add_column("x".into(), vec![Value::Float(2.0)]);
438
439        df1.vstack(&df2);
440        assert_eq!(df1.nrows(), 2);
441    }
442
443    #[test]
444    fn test_sort_by() {
445        let mut df = DataFrame::new();
446        df.add_column(
447            "x".into(),
448            vec![Value::Float(3.0), Value::Float(1.0), Value::Float(2.0)],
449        );
450        let sorted = df.sort_by("x");
451        let col = sorted.column("x").unwrap();
452        assert_eq!(col[0].as_f64(), Some(1.0));
453        assert_eq!(col[1].as_f64(), Some(2.0));
454        assert_eq!(col[2].as_f64(), Some(3.0));
455    }
456}