use std::collections::HashMap;
use std::io::{Cursor, Write};
use std::sync::Arc;
use arrow_array::RecordBatch;
use arrow_schema::{ArrowError, Field, Schema};
use bytes::Bytes;
use parquet::arrow::ArrowWriter;
use parquet::arrow::PARQUET_FIELD_ID_META_KEY;
use parquet::basic::Compression;
use parquet::file::properties::WriterProperties;
use crate::error::Error;
fn add_field_ids_to_schema(schema: &Schema) -> Schema {
let fields_with_ids: Vec<Arc<Field>> = schema
.fields()
.iter()
.enumerate()
.map(|(idx, field)| {
let field_id = (idx + 1) as i32;
let mut metadata: HashMap<String, String> = field.metadata().clone();
metadata.insert(PARQUET_FIELD_ID_META_KEY.to_string(), field_id.to_string());
Arc::new(field.as_ref().clone().with_metadata(metadata))
})
.collect();
Schema::new_with_metadata(fields_with_ids, schema.metadata().clone())
}
pub fn write_parquet<W: Write + Send>(
batch: &RecordBatch,
writer: W,
props: Option<WriterProperties>,
) -> Result<(), Error> {
let props = props.unwrap_or_else(|| {
WriterProperties::builder()
.set_compression(Compression::UNCOMPRESSED)
.build()
});
let schema_with_ids = Arc::new(add_field_ids_to_schema(batch.schema().as_ref()));
let batch_with_ids = RecordBatch::try_new(schema_with_ids.clone(), batch.columns().to_vec())
.map_err(Error::Arrow)?;
let mut arrow_writer = ArrowWriter::try_new(writer, schema_with_ids, Some(props))
.map_err(|e| Error::Arrow(ArrowError::ExternalError(Box::new(e))))?;
arrow_writer
.write(&batch_with_ids)
.map_err(|e| Error::Arrow(ArrowError::ExternalError(Box::new(e))))?;
arrow_writer
.close()
.map_err(|e| Error::Arrow(ArrowError::ExternalError(Box::new(e))))?;
Ok(())
}
pub fn to_parquet(batch: &RecordBatch) -> Result<Vec<u8>, Error> {
let mut buffer = Cursor::new(Vec::new());
write_parquet(batch, &mut buffer, None)?;
Ok(buffer.into_inner())
}
#[allow(dead_code)]
pub fn to_parquet_bytes(batch: &RecordBatch) -> Result<Bytes, Error> {
let vec = to_parquet(batch)?;
Ok(Bytes::from(vec))
}
#[cfg(test)]
mod tests {
use super::*;
use arrow_array::{Array, Int64Array, StringArray};
use arrow_schema::{DataType, Field, Schema};
use parquet::arrow::arrow_reader::ParquetRecordBatchReaderBuilder;
use std::sync::Arc;
fn create_test_batch() -> RecordBatch {
let schema = Arc::new(Schema::new(vec![
Field::new("name", DataType::Utf8, false),
Field::new("value", DataType::Int64, false),
]));
let name_array = Arc::new(StringArray::from(vec!["alpha", "beta", "gamma"]));
let value_array = Arc::new(Int64Array::from(vec![1, 2, 3]));
RecordBatch::try_new(schema, vec![name_array, value_array]).unwrap()
}
#[test]
fn test_to_parquet_basic() {
let batch = create_test_batch();
let result = to_parquet(&batch).unwrap();
assert!(!result.is_empty());
assert_eq!(&result[0..4], b"PAR1");
}
#[test]
fn test_to_parquet_roundtrip() {
let original_batch = create_test_batch();
let parquet_bytes = to_parquet(&original_batch).unwrap();
let reader = ParquetRecordBatchReaderBuilder::try_new(Bytes::from(parquet_bytes))
.unwrap()
.build()
.unwrap();
let batches: Vec<RecordBatch> = reader.map(|r| r.unwrap()).collect();
assert_eq!(batches.len(), 1);
let read_batch = &batches[0];
assert_eq!(read_batch.num_rows(), 3);
assert_eq!(read_batch.num_columns(), 2);
let name_col = read_batch
.column(0)
.as_any()
.downcast_ref::<StringArray>()
.unwrap();
let value_col = read_batch
.column(1)
.as_any()
.downcast_ref::<Int64Array>()
.unwrap();
assert_eq!(name_col.value(0), "alpha");
assert_eq!(name_col.value(1), "beta");
assert_eq!(name_col.value(2), "gamma");
assert_eq!(value_col.value(0), 1);
assert_eq!(value_col.value(1), 2);
assert_eq!(value_col.value(2), 3);
}
#[test]
fn test_to_parquet_empty_batch() {
let schema = Arc::new(Schema::new(vec![
Field::new("name", DataType::Utf8, false),
Field::new("value", DataType::Int64, false),
]));
let name_array = Arc::new(StringArray::from(Vec::<&str>::new()));
let value_array = Arc::new(Int64Array::from(Vec::<i64>::new()));
let batch = RecordBatch::try_new(schema, vec![name_array, value_array]).unwrap();
let result = to_parquet(&batch).unwrap();
assert!(!result.is_empty());
assert_eq!(&result[0..4], b"PAR1");
let reader = ParquetRecordBatchReaderBuilder::try_new(Bytes::from(result))
.unwrap()
.build()
.unwrap();
let batches: Vec<RecordBatch> = reader.map(|r| r.unwrap()).collect();
let total_rows: usize = batches.iter().map(|b| b.num_rows()).sum();
assert_eq!(total_rows, 0);
}
#[test]
fn test_to_parquet_with_nulls() {
let schema = Arc::new(Schema::new(vec![
Field::new("name", DataType::Utf8, true),
Field::new("value", DataType::Int64, true),
]));
let name_array = Arc::new(StringArray::from(vec![Some("alpha"), None, Some("gamma")]));
let value_array = Arc::new(Int64Array::from(vec![Some(1), Some(2), None]));
let batch = RecordBatch::try_new(schema, vec![name_array, value_array]).unwrap();
let result = to_parquet(&batch).unwrap();
let reader = ParquetRecordBatchReaderBuilder::try_new(Bytes::from(result))
.unwrap()
.build()
.unwrap();
let batches: Vec<RecordBatch> = reader.map(|r| r.unwrap()).collect();
let read_batch = &batches[0];
let name_col = read_batch
.column(0)
.as_any()
.downcast_ref::<StringArray>()
.unwrap();
let value_col = read_batch
.column(1)
.as_any()
.downcast_ref::<Int64Array>()
.unwrap();
assert!(!name_col.is_null(0));
assert!(name_col.is_null(1));
assert!(!name_col.is_null(2));
assert!(!value_col.is_null(0));
assert!(!value_col.is_null(1));
assert!(value_col.is_null(2));
}
#[test]
fn test_to_parquet_bytes() {
let batch = create_test_batch();
let result = to_parquet_bytes(&batch).unwrap();
assert!(!result.is_empty());
assert_eq!(&result[0..4], b"PAR1");
}
#[test]
fn test_write_parquet_to_cursor() {
let batch = create_test_batch();
let mut buffer = Cursor::new(Vec::new());
write_parquet(&batch, &mut buffer, None).unwrap();
let result = buffer.into_inner();
assert!(!result.is_empty());
assert_eq!(&result[0..4], b"PAR1");
let reader = ParquetRecordBatchReaderBuilder::try_new(Bytes::from(result))
.unwrap()
.build()
.unwrap();
let batches: Vec<RecordBatch> = reader.map(|r| r.unwrap()).collect();
assert_eq!(batches.len(), 1);
assert_eq!(batches[0].num_rows(), 3);
}
#[test]
fn test_write_parquet_with_custom_properties() {
let batch = create_test_batch();
let mut buffer = Cursor::new(Vec::new());
let props = WriterProperties::builder()
.set_compression(Compression::UNCOMPRESSED)
.set_data_page_row_count_limit(100)
.build();
write_parquet(&batch, &mut buffer, Some(props)).unwrap();
let result = buffer.into_inner();
assert!(!result.is_empty());
assert_eq!(&result[0..4], b"PAR1");
let reader = ParquetRecordBatchReaderBuilder::try_new(Bytes::from(result))
.unwrap()
.build()
.unwrap();
let batches: Vec<RecordBatch> = reader.map(|r| r.unwrap()).collect();
assert_eq!(batches.len(), 1);
assert_eq!(batches[0].num_rows(), 3);
}
}