alopex_dataframe/dataframe/
dataframe.rs1use 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#[derive(Debug, Clone)]
12pub struct DataFrame {
13 schema: SchemaRef,
14 batches: Vec<RecordBatch>,
15}
16
17impl DataFrame {
18 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 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 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 pub fn from_series(series: Vec<Series>) -> Result<Self> {
115 Self::new(series)
116 }
117
118 pub fn empty() -> Self {
120 Self {
121 schema: Arc::new(Schema::empty()),
122 batches: Vec::new(),
123 }
124 }
125
126 pub fn height(&self) -> usize {
128 self.batches.iter().map(|b| b.num_rows()).sum()
129 }
130
131 pub fn width(&self) -> usize {
133 self.schema.fields().len()
134 }
135
136 pub fn schema(&self) -> SchemaRef {
138 self.schema.clone()
139 }
140
141 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 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 pub fn to_arrow(&self) -> Vec<RecordBatch> {
177 self.batches.clone()
178 }
179
180 pub fn lazy(&self) -> crate::LazyFrame {
182 crate::LazyFrame::from_dataframe(self.clone())
183 }
184
185 pub fn select(&self, exprs: Vec<Expr>) -> Result<Self> {
187 self.clone().lazy().select(exprs).collect()
188 }
189
190 pub fn filter(&self, predicate: Expr) -> Result<Self> {
192 self.clone().lazy().filter(predicate).collect()
193 }
194
195 pub fn with_columns(&self, exprs: Vec<Expr>) -> Result<Self> {
197 self.clone().lazy().with_columns(exprs).collect()
198 }
199
200 pub fn group_by(&self, by: Vec<Expr>) -> GroupBy {
202 GroupBy {
203 df: self.clone(),
204 by,
205 }
206 }
207
208 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 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 pub fn head(&self, n: usize) -> Result<Self> {
234 self.clone().lazy().head(n).collect()
235 }
236
237 pub fn tail(&self, n: usize) -> Result<Self> {
239 self.clone().lazy().tail(n).collect()
240 }
241
242 pub fn unique(&self, subset: Option<Vec<String>>) -> Result<Self> {
244 self.clone().lazy().unique(subset).collect()
245 }
246
247 pub fn fill_null<T: Into<FillNull>>(&self, fill: T) -> Result<Self> {
249 self.clone().lazy().fill_null(fill).collect()
250 }
251
252 pub fn drop_nulls(&self, subset: Option<Vec<String>>) -> Result<Self> {
254 self.clone().lazy().drop_nulls(subset).collect()
255 }
256
257 pub fn null_count(&self) -> Result<Self> {
259 self.clone().lazy().null_count().collect()
260 }
261
262 pub fn explode(&self, column: impl Into<String>) -> Result<Self> {
264 self.clone().lazy().explode(column).collect()
265 }
266
267 pub fn implode(&self) -> Result<Self> {
269 self.clone().lazy().implode().collect()
270 }
271}
272
273#[derive(Debug, Clone)]
275pub struct GroupBy {
276 df: DataFrame,
277 by: Vec<Expr>,
278}
279
280impl GroupBy {
281 pub fn agg(self, aggs: Vec<Expr>) -> Result<DataFrame> {
283 self.df.lazy().group_by(self.by).agg(aggs).collect()
284 }
285
286 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}