Skip to main content

alopex_dataframe/dataframe/
dataframe.rs

1use std::collections::HashSet;
2use std::sync::Arc;
3
4use arrow::datatypes::{Field, Schema, SchemaRef};
5use arrow::record_batch::RecordBatch;
6
7use crate::ops::{FillNull, JoinKeys, JoinType, SortOptions};
8use crate::{DataFrameError, Expr, Result, Series};
9
10/// An eager table backed by one or more Arrow `RecordBatch` values.
11#[derive(Debug, Clone)]
12pub struct DataFrame {
13    schema: SchemaRef,
14    batches: Vec<RecordBatch>,
15}
16
17impl DataFrame {
18    /// Construct a `DataFrame` from a list of `Series`.
19    ///
20    /// Chunk boundaries do not need to align across series as long as total lengths match.
21    pub fn new(columns: Vec<Series>) -> Result<Self> {
22        if columns.is_empty() {
23            return Ok(Self::empty());
24        }
25
26        let mut seen_names = HashSet::with_capacity(columns.len());
27        for c in &columns {
28            if !seen_names.insert(c.name().to_string()) {
29                return Err(DataFrameError::schema_mismatch(format!(
30                    "duplicate column name '{}'",
31                    c.name()
32                )));
33            }
34        }
35
36        let expected_len = columns[0].len();
37        for c in &columns[1..] {
38            if c.len() != expected_len {
39                return Err(DataFrameError::schema_mismatch(format!(
40                    "column length mismatch: '{}' has length {}, expected {}",
41                    c.name(),
42                    c.len(),
43                    expected_len
44                )));
45            }
46        }
47
48        let fields: Vec<Field> = columns
49            .iter()
50            .map(|c| Field::new(c.name(), c.dtype(), true))
51            .collect();
52        let schema: SchemaRef = Arc::new(Schema::new(fields));
53
54        let arrays = columns
55            .iter()
56            .map(|c| {
57                if c.chunks().is_empty() {
58                    Ok(arrow::array::new_empty_array(&c.dtype()))
59                } else if c.chunks().len() == 1 {
60                    Ok(c.chunks()[0].clone())
61                } else {
62                    let arrays = c
63                        .chunks()
64                        .iter()
65                        .map(|a| a.as_ref() as &dyn arrow::array::Array)
66                        .collect::<Vec<_>>();
67                    arrow::compute::concat(&arrays)
68                        .map_err(|source| DataFrameError::Arrow { source })
69                }
70            })
71            .collect::<Result<Vec<_>>>()?;
72
73        let batch = RecordBatch::try_new(schema.clone(), arrays).map_err(|e| {
74            DataFrameError::schema_mismatch(format!("failed to build RecordBatch: {e}"))
75        })?;
76
77        Ok(Self {
78            schema,
79            batches: vec![batch],
80        })
81    }
82
83    /// Construct a `DataFrame` from Arrow record batches (all batches must share the same schema).
84    pub fn from_batches(batches: Vec<RecordBatch>) -> Result<Self> {
85        if batches.is_empty() {
86            return Ok(Self::empty());
87        }
88
89        let schema = batches[0].schema();
90        for (i, b) in batches.iter().enumerate().skip(1) {
91            if b.schema().as_ref() != schema.as_ref() {
92                return Err(DataFrameError::schema_mismatch(format!(
93                    "schema mismatch between batches: batch 0 != batch {i}"
94                )));
95            }
96        }
97
98        Ok(Self { schema, batches })
99    }
100
101    /// Strict vertical concatenation of two or more eager frames.
102    ///
103    /// Column name/order/type/nullability must match exactly. The result keeps every batch from
104    /// input one before every batch from input two, and so on; no implicit coercion is performed.
105    pub fn concat(inputs: Vec<DataFrame>) -> Result<Self> {
106        let lazy_inputs = inputs
107            .into_iter()
108            .map(crate::LazyFrame::from_dataframe)
109            .collect();
110        crate::LazyFrame::concat(lazy_inputs)?.collect()
111    }
112
113    /// Alias for `DataFrame::new`.
114    pub fn from_series(series: Vec<Series>) -> Result<Self> {
115        Self::new(series)
116    }
117
118    /// Return an empty `DataFrame` (no columns, no rows).
119    pub fn empty() -> Self {
120        Self {
121            schema: Arc::new(Schema::empty()),
122            batches: Vec::new(),
123        }
124    }
125
126    /// Return the number of rows.
127    pub fn height(&self) -> usize {
128        self.batches.iter().map(|b| b.num_rows()).sum()
129    }
130
131    /// Return the number of columns.
132    pub fn width(&self) -> usize {
133        self.schema.fields().len()
134    }
135
136    /// Return the Arrow schema.
137    pub fn schema(&self) -> SchemaRef {
138        self.schema.clone()
139    }
140
141    /// Get a column by name (case-sensitive).
142    pub fn column(&self, name: &str) -> Result<Series> {
143        let idx = self
144            .schema
145            .fields()
146            .iter()
147            .position(|f| f.name() == name)
148            .ok_or_else(|| DataFrameError::column_not_found(name.to_string()))?;
149
150        let chunks = self
151            .batches
152            .iter()
153            .map(|b| b.column(idx).clone())
154            .collect::<Vec<_>>();
155        Ok(Series::from_arrow_unchecked(name, chunks))
156    }
157
158    /// Return all columns in construction order.
159    pub fn columns(&self) -> Vec<Series> {
160        self.schema
161            .fields()
162            .iter()
163            .enumerate()
164            .map(|(idx, f)| {
165                let chunks = self
166                    .batches
167                    .iter()
168                    .map(|b| b.column(idx).clone())
169                    .collect::<Vec<_>>();
170                Series::from_arrow_unchecked(f.name(), chunks)
171            })
172            .collect()
173    }
174
175    /// Return the underlying Arrow batches.
176    pub fn to_arrow(&self) -> Vec<RecordBatch> {
177        self.batches.clone()
178    }
179
180    /// Convert this eager `DataFrame` to a `LazyFrame` for query planning/execution.
181    pub fn lazy(&self) -> crate::LazyFrame {
182        crate::LazyFrame::from_dataframe(self.clone())
183    }
184
185    /// Eager `select`, implemented by delegating to `LazyFrame`.
186    pub fn select(&self, exprs: Vec<Expr>) -> Result<Self> {
187        self.clone().lazy().select(exprs).collect()
188    }
189
190    /// Eager `filter`, implemented by delegating to `LazyFrame`.
191    pub fn filter(&self, predicate: Expr) -> Result<Self> {
192        self.clone().lazy().filter(predicate).collect()
193    }
194
195    /// Eager `with_columns`, implemented by delegating to `LazyFrame`.
196    pub fn with_columns(&self, exprs: Vec<Expr>) -> Result<Self> {
197        self.clone().lazy().with_columns(exprs).collect()
198    }
199
200    /// Start a group-by aggregation (eager API).
201    pub fn group_by(&self, by: Vec<Expr>) -> GroupBy {
202        GroupBy {
203            df: self.clone(),
204            by,
205        }
206    }
207
208    /// Join with another `DataFrame` using provided join keys.
209    pub fn join<K: Into<JoinKeys>>(
210        &self,
211        other: &DataFrame,
212        keys: K,
213        how: JoinType,
214    ) -> Result<Self> {
215        self.clone()
216            .lazy()
217            .join(other.clone().lazy(), keys, how)
218            .collect()
219    }
220
221    /// Sort by one or more columns.
222    pub fn sort(&self, by: Vec<String>, descending: Vec<bool>) -> Result<Self> {
223        let options = SortOptions {
224            by,
225            descending,
226            nulls_last: true,
227            stable: true,
228        };
229        self.clone().lazy().sort(options).collect()
230    }
231
232    /// Return the first `n` rows.
233    pub fn head(&self, n: usize) -> Result<Self> {
234        self.clone().lazy().head(n).collect()
235    }
236
237    /// Return the last `n` rows.
238    pub fn tail(&self, n: usize) -> Result<Self> {
239        self.clone().lazy().tail(n).collect()
240    }
241
242    /// Remove duplicate rows.
243    pub fn unique(&self, subset: Option<Vec<String>>) -> Result<Self> {
244        self.clone().lazy().unique(subset).collect()
245    }
246
247    /// Fill null values using a scalar or strategy.
248    pub fn fill_null<T: Into<FillNull>>(&self, fill: T) -> Result<Self> {
249        self.clone().lazy().fill_null(fill).collect()
250    }
251
252    /// Drop rows containing null values.
253    pub fn drop_nulls(&self, subset: Option<Vec<String>>) -> Result<Self> {
254        self.clone().lazy().drop_nulls(subset).collect()
255    }
256
257    /// Count null values per column.
258    pub fn null_count(&self) -> Result<Self> {
259        self.clone().lazy().null_count().collect()
260    }
261
262    /// Explode one `List<Utf8>` column into multiple rows.
263    pub fn explode(&self, column: impl Into<String>) -> Result<Self> {
264        self.clone().lazy().explode(column).collect()
265    }
266
267    /// Implode UTF-8 columns into one row of `List<Utf8>` columns.
268    pub fn implode(&self) -> Result<Self> {
269        self.clone().lazy().implode().collect()
270    }
271}
272
273/// Eager group-by handle that delegates execution to `LazyFrame`.
274#[derive(Debug, Clone)]
275pub struct GroupBy {
276    df: DataFrame,
277    by: Vec<Expr>,
278}
279
280impl GroupBy {
281    /// Perform aggregations for this group-by.
282    pub fn agg(self, aggs: Vec<Expr>) -> Result<DataFrame> {
283        self.df.lazy().group_by(self.by).agg(aggs).collect()
284    }
285
286    /// Return the underlying `DataFrame`.
287    pub fn into_df(self) -> DataFrame {
288        self.df
289    }
290}
291
292#[cfg(test)]
293mod tests {
294    use std::sync::Arc;
295
296    use arrow::array::{ArrayRef, Int32Array, StringArray};
297    use arrow::datatypes::{DataType, Field, Schema};
298    use arrow::record_batch::RecordBatch;
299
300    use super::DataFrame;
301    use crate::{DataFrameError, Series};
302
303    fn s_i32(name: &str, chunks: Vec<Vec<i32>>) -> Series {
304        let arrays: Vec<ArrayRef> = chunks
305            .into_iter()
306            .map(|v| Arc::new(Int32Array::from(v)) as ArrayRef)
307            .collect();
308        Series::from_arrow(name, arrays).unwrap()
309    }
310
311    #[test]
312    fn dataframe_new_accepts_misaligned_chunks_by_normalizing() {
313        let a = s_i32("a", vec![vec![1, 2], vec![3]]);
314        let b = s_i32("b", vec![vec![10], vec![20, 30]]);
315
316        let df = DataFrame::new(vec![a, b]).unwrap();
317        assert_eq!(df.height(), 3);
318        assert_eq!(df.width(), 2);
319        assert_eq!(df.schema().fields()[0].name(), "a");
320        assert_eq!(df.schema().fields()[1].name(), "b");
321
322        let batches = df.to_arrow();
323        assert_eq!(batches.len(), 1);
324        assert_eq!(batches[0].num_rows(), 3);
325    }
326
327    #[test]
328    fn dataframe_new_rejects_duplicate_column_names() {
329        let a1 = s_i32("a", vec![vec![1]]);
330        let a2 = s_i32("a", vec![vec![2]]);
331        let err = DataFrame::new(vec![a1, a2]).unwrap_err();
332        assert!(matches!(err, DataFrameError::SchemaMismatch { .. }));
333    }
334
335    #[test]
336    fn dataframe_new_rejects_length_mismatch() {
337        let a = s_i32("a", vec![vec![1, 2]]);
338        let b = s_i32("b", vec![vec![10]]);
339        let err = DataFrame::new(vec![a, b]).unwrap_err();
340        assert!(matches!(err, DataFrameError::SchemaMismatch { .. }));
341    }
342
343    #[test]
344    fn dataframe_new_accepts_different_chunk_counts() {
345        let a = s_i32("a", vec![vec![1], vec![2], vec![3]]);
346        let b = s_i32("b", vec![vec![10, 20, 30]]);
347        let df = DataFrame::new(vec![a, b]).unwrap();
348        assert_eq!(df.height(), 3);
349        assert_eq!(df.to_arrow().len(), 1);
350    }
351
352    #[test]
353    fn dataframe_column_is_case_sensitive() {
354        let a = s_i32("a", vec![vec![1]]);
355        let df = DataFrame::new(vec![a]).unwrap();
356        assert!(matches!(
357            df.column("A").unwrap_err(),
358            DataFrameError::ColumnNotFound { .. }
359        ));
360    }
361
362    #[test]
363    fn dataframe_from_batches_rejects_schema_mismatch() {
364        let a1: ArrayRef = Arc::new(Int32Array::from(vec![1]));
365        let a2: ArrayRef = Arc::new(StringArray::from(vec!["x"]));
366
367        let s1 = Arc::new(Schema::new(vec![Field::new("a", DataType::Int32, true)]));
368        let s2 = Arc::new(Schema::new(vec![Field::new("a", DataType::Utf8, true)]));
369
370        let b1 = RecordBatch::try_new(s1, vec![a1]).unwrap();
371        let b2 = RecordBatch::try_new(s2, vec![a2]).unwrap();
372
373        let err = DataFrame::from_batches(vec![b1, b2]).unwrap_err();
374        assert!(matches!(err, DataFrameError::SchemaMismatch { .. }));
375    }
376
377    #[test]
378    fn dataframe_columns_preserves_schema_order() {
379        let a = s_i32("a", vec![vec![1], vec![2]]);
380        let b = s_i32("b", vec![vec![10], vec![20]]);
381        let df = DataFrame::new(vec![b.clone(), a.clone()]).unwrap();
382
383        let cols = df.columns();
384        assert_eq!(cols[0].name(), "b");
385        assert_eq!(cols[1].name(), "a");
386        assert_eq!(cols[0].len(), 2);
387        assert_eq!(cols[1].len(), 2);
388    }
389
390    #[test]
391    fn eager_concat_preserves_input_order_and_rejects_mismatched_schema() {
392        let first = DataFrame::new(vec![Series::from_arrow(
393            "value",
394            vec![Arc::new(Int32Array::from(vec![1])) as ArrayRef],
395        )
396        .unwrap()])
397        .unwrap();
398        let second = DataFrame::new(vec![Series::from_arrow(
399            "value",
400            vec![Arc::new(Int32Array::from(vec![2])) as ArrayRef],
401        )
402        .unwrap()])
403        .unwrap();
404        let output = DataFrame::concat(vec![first.clone(), second]).unwrap();
405        let values = output
406            .to_arrow()
407            .into_iter()
408            .flat_map(|batch| {
409                batch
410                    .column(0)
411                    .as_any()
412                    .downcast_ref::<Int32Array>()
413                    .unwrap()
414                    .values()
415                    .iter()
416                    .copied()
417                    .collect::<Vec<_>>()
418            })
419            .collect::<Vec<_>>();
420        assert_eq!(values, vec![1, 2]);
421
422        let incompatible = DataFrame::new(vec![Series::from_arrow(
423            "other",
424            vec![Arc::new(Int32Array::from(vec![3])) as ArrayRef],
425        )
426        .unwrap()])
427        .unwrap();
428        let err = DataFrame::concat(vec![first, incompatible]).unwrap_err();
429        assert!(err.to_string().contains("concat_schema_mismatch"));
430    }
431}