1use indexmap::IndexMap;
2
3use super::{GroupKey, Value};
4
5#[derive(Clone, Debug)]
7pub struct DataFrame {
8 columns: IndexMap<String, Vec<Value>>,
9 nrows: usize,
10 issues: Vec<String>,
14}
15
16impl DataFrame {
17 pub fn new() -> Self {
19 DataFrame {
20 columns: IndexMap::new(),
21 nrows: 0,
22 issues: Vec::new(),
23 }
24 }
25
26 pub fn column(&self, name: &str) -> Option<&[Value]> {
28 self.columns.get(name).map(|v| v.as_slice())
29 }
30
31 pub fn nrows(&self) -> usize {
33 self.nrows
34 }
35
36 pub fn ncols(&self) -> usize {
38 self.columns.len()
39 }
40
41 pub fn column_names(&self) -> Vec<&str> {
43 self.columns.keys().map(|s| s.as_str()).collect()
44 }
45
46 pub fn has_column(&self, name: &str) -> bool {
48 self.columns.contains_key(name)
49 }
50
51 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 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 pub fn issues(&self) -> &[String] {
99 &self.issues
100 }
101
102 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 pub fn column_mut(&mut self, name: &str) -> Option<&mut Vec<Value>> {
116 self.columns.get_mut(name)
117 }
118
119 pub fn group_by(&self, keys: &[&str]) -> Vec<DataFrame> {
121 if self.nrows == 0 {
122 return vec![];
123 }
124
125 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 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 for (name, col) in &self.columns {
165 if let Some(other_col) = other.columns.get(name) {
166 let _ = (col, other_col);
168 }
169 }
170
171 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 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 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 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 pub fn from_rows(rows: Vec<IndexMap<String, Value>>) -> Self {
241 if rows.is_empty() {
242 return DataFrame::new();
243 }
244
245 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 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 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 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); assert_eq!(groups[1].nrows(), 1); }
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}