use arrow::array::RecordBatch as ArrowRecordBatch;
use arrow::datatypes::Schema as ArrowSchema;
use arrow::ipc::reader::StreamReader;
use polars::error::PolarsError;
use polars::frame::DataFrame;
use polars::prelude::*;
use polars_io::ipc::{IpcCompression, IpcStreamWriter};
use std::io::Cursor;
use std::sync::Arc;
pub fn write_to_arrow(
df: &mut DataFrame,
compression: Option<IpcCompression>,
) -> PolarsResult<(Arc<ArrowSchema>, Vec<ArrowRecordBatch>)> {
let mut buffer = Vec::new();
IpcStreamWriter::new(&mut buffer)
.with_compression(compression)
.finish(df)?;
let cursor = Cursor::new(buffer);
let stream_reader = StreamReader::try_new(cursor, None)
.map_err(|e| PolarsError::ComputeError(format!("{e}").into()))?;
let schema = stream_reader.schema();
let batches: Result<Vec<_>, _> = stream_reader.collect();
let batches = batches.map_err(|e| PolarsError::ComputeError(format!("{e}").into()))?;
Ok((schema, batches))
}
#[cfg(test)]
mod tests {
use super::*;
fn create_df() -> DataFrame {
df!(
"ticker" => ["AAPL", "NVDA", "MSFT", "GOOG", "AMZN"],
"price" => [229.9, 138.93, 420.56, 166.41, 188.4],
"high" => [231.31, 139.6, 424.04, 167.62, 189.83],
"low" => [228.6, 136.3, 417.52, 164.78, 188.44],
)
.unwrap()
}
#[test]
fn compressed_write() {
let mut df = create_df();
let arrow_data = write_to_arrow(&mut df, Some(IpcCompression::ZSTD));
let (arrow_schema, record_batches) = arrow_data.unwrap();
let df_schema = df.schema();
let arrow_schema_fields = arrow_schema.fields();
let df_schema_fields: Vec<Field> = df_schema.iter_fields().collect();
let column_names = df.get_column_names();
let column_names_owned = df.get_column_names_owned();
assert_eq!(arrow_schema_fields.len(), df_schema_fields.len());
assert_eq!(arrow_schema_fields.len(), df.width());
for (i, field) in arrow_schema_fields.iter().enumerate() {
assert_eq!(field.name().to_owned(), column_names_owned[i].to_string());
}
assert!(!record_batches.is_empty());
assert_eq!(record_batches.len(), 1);
let batch = &record_batches[0];
assert_eq!(batch.num_rows(), df.height());
assert_eq!(batch.num_columns(), df.width());
assert_eq!(column_names.len(), batch.num_columns());
assert_eq!(column_names_owned.len(), batch.num_columns());
}
#[test]
fn uncompressed_write() {
let mut df = create_df();
let arrow_data = write_to_arrow(&mut df, None);
let (arrow_schema, record_batches) = arrow_data.unwrap();
let df_schema = df.schema();
let arrow_schema_fields = arrow_schema.fields();
let df_schema_fields: Vec<Field> = df_schema.iter_fields().collect();
let column_names = df.get_column_names();
let column_names_owned = df.get_column_names_owned();
assert_eq!(arrow_schema_fields.len(), df_schema_fields.len());
assert_eq!(arrow_schema_fields.len(), df.width());
for (i, field) in arrow_schema_fields.iter().enumerate() {
assert_eq!(field.name().to_owned(), column_names_owned[i].to_string());
}
assert!(!record_batches.is_empty());
assert_eq!(record_batches.len(), 1);
let batch = &record_batches[0];
assert_eq!(batch.num_rows(), df.height());
assert_eq!(batch.num_columns(), df.width());
assert_eq!(column_names.len(), batch.num_columns());
assert_eq!(column_names_owned.len(), batch.num_columns());
}
}