use minarrow::{Array, Bitmask, Integer, NumericArray, SuperTable, Table, TextArray};
use std::io::{self, Write};
#[derive(Debug, Clone)]
pub struct CsvEncodeOptions {
pub delimiter: u8,
pub write_header: bool,
pub null_repr: &'static str,
pub quote: u8,
}
impl Default for CsvEncodeOptions {
fn default() -> Self {
CsvEncodeOptions {
delimiter: b',',
write_header: true,
null_repr: "",
quote: b'"',
}
}
}
#[inline]
fn needs_quoting(bytes: &[u8], delimiter: u8, quote: u8) -> bool {
if bytes.first() == Some(&b' ') || bytes.last() == Some(&b' ') {
return true;
}
memchr::memchr3(delimiter, quote, b'\n', bytes).is_some()
|| memchr::memchr(b'\r', bytes).is_some()
}
#[inline]
fn append_csv_string(out: &mut Vec<u8>, bytes: &[u8], delimiter: u8, quote: u8) {
if !needs_quoting(bytes, delimiter, quote) {
out.extend_from_slice(bytes);
return;
}
out.push(quote);
let mut start = 0;
while let Some(pos) = memchr::memchr(quote, &bytes[start..]) {
let abs = start + pos;
out.extend_from_slice(&bytes[start..abs]);
out.push(quote);
out.push(quote);
start = abs + 1;
}
out.extend_from_slice(&bytes[start..]);
out.push(quote);
}
#[inline]
fn append_str_cell<T: Integer>(
out: &mut Vec<u8>,
data: &[u8],
offsets: &[T],
row: usize,
delimiter: u8,
quote: u8,
) {
let start = offsets[row].to_usize();
let mut end = offsets[row + 1].to_usize();
if row + 1 == offsets.len() - 1 && end < data.len() {
end = data.len();
}
append_csv_string(out, &data[start..end], delimiter, quote);
}
fn estimate_cell_width(arr: &Array, n_rows: usize) -> usize {
match arr {
Array::NumericArray(n) => match n {
NumericArray::Int32(_) | NumericArray::UInt32(_) => 6,
NumericArray::Int64(_) | NumericArray::UInt64(_) => 10,
#[cfg(feature = "extended_numeric_types")]
NumericArray::Int8(_) | NumericArray::UInt8(_) => 4,
#[cfg(feature = "extended_numeric_types")]
NumericArray::Int16(_) | NumericArray::UInt16(_) => 5,
NumericArray::Float32(_) | NumericArray::Float64(_) => 16,
_ => 12,
},
Array::BooleanArray(_) => 5,
Array::TextArray(TextArray::String32(arr)) => {
arr.data.len().checked_div(n_rows).map_or(8, |v| v + 4)
}
#[cfg(feature = "large_string")]
Array::TextArray(TextArray::String64(arr)) => {
if n_rows == 0 {
8
} else {
arr.data.len() / n_rows + 4
}
}
Array::TextArray(_) => {
12
}
#[cfg(feature = "datetime")]
Array::TemporalArray(_) => 12,
_ => 12,
}
}
pub fn encode_table_csv<W: Write>(
table: &Table,
mut writer: W,
options: &CsvEncodeOptions,
) -> io::Result<()> {
let CsvEncodeOptions {
delimiter,
write_header,
null_repr,
quote,
} = *options;
let n_rows = table.n_rows;
let mut null_masks: Vec<Option<&Bitmask>> = Vec::with_capacity(table.cols.len());
for col in &table.cols {
match &col.array {
Array::NumericArray(arr) => null_masks.push(arr.null_mask()),
Array::BooleanArray(arr) => null_masks.push(arr.null_mask.as_ref()),
Array::TextArray(TextArray::String32(arr)) => null_masks.push(arr.null_mask.as_ref()),
#[cfg(any(
not(feature = "default_categorical_8"),
feature = "extended_categorical"
))]
Array::TextArray(TextArray::Categorical32(arr)) => {
null_masks.push(arr.null_mask.as_ref())
}
#[cfg(feature = "large_string")]
Array::TextArray(TextArray::String64(arr)) => null_masks.push(arr.null_mask.as_ref()),
#[cfg(feature = "default_categorical_8")]
Array::TextArray(TextArray::Categorical8(arr)) => {
null_masks.push(arr.null_mask.as_ref())
}
#[cfg(feature = "extended_categorical")]
Array::TextArray(TextArray::Categorical16(arr)) => {
null_masks.push(arr.null_mask.as_ref())
}
#[cfg(feature = "extended_categorical")]
Array::TextArray(TextArray::Categorical64(arr)) => {
null_masks.push(arr.null_mask.as_ref())
}
#[cfg(feature = "datetime")]
Array::TemporalArray(arr) => {
let null_mask = match arr {
minarrow::TemporalArray::Datetime32(arr) => arr.null_mask.as_ref(),
minarrow::TemporalArray::Datetime64(arr) => arr.null_mask.as_ref(),
minarrow::TemporalArray::Null => None,
};
null_masks.push(null_mask)
}
_ => null_masks.push(None),
}
}
let est_row: usize = table
.cols
.iter()
.map(|c| estimate_cell_width(&c.array, n_rows) + 1)
.sum::<usize>();
let header_est = if write_header {
table
.cols
.iter()
.map(|c| c.field.name.len() + 1)
.sum::<usize>()
} else {
0
};
let mut out: Vec<u8> = Vec::with_capacity(header_est + est_row * n_rows + 16);
if write_header {
for (i, col) in table.cols.iter().enumerate() {
if i > 0 {
out.push(delimiter);
}
append_csv_string(&mut out, col.field.name.as_bytes(), delimiter, quote);
}
out.push(b'\n');
}
let null_bytes = null_repr.as_bytes();
let mut itoa_buf = super::int_ascii::Buffer::new();
let mut ryu_buf = ryu::Buffer::new();
for row in 0..n_rows {
for (col_idx, col) in table.cols.iter().enumerate() {
if col_idx > 0 {
out.push(delimiter);
}
let is_null = if col.null_count == 0 {
false
} else {
match null_masks[col_idx] {
Some(mask) => !mask.get(row),
None => false,
}
};
if is_null {
out.extend_from_slice(null_bytes);
continue;
}
match &col.array {
Array::NumericArray(n) => match n {
#[cfg(feature = "extended_numeric_types")]
NumericArray::Int8(arr) => {
out.extend_from_slice(itoa_buf.format(arr.data.as_ref()[row]).as_bytes());
}
#[cfg(feature = "extended_numeric_types")]
NumericArray::Int16(arr) => {
out.extend_from_slice(itoa_buf.format(arr.data.as_ref()[row]).as_bytes());
}
NumericArray::Int32(arr) => {
out.extend_from_slice(itoa_buf.format(arr.data.as_ref()[row]).as_bytes());
}
NumericArray::Int64(arr) => {
out.extend_from_slice(itoa_buf.format(arr.data.as_ref()[row]).as_bytes());
}
#[cfg(feature = "extended_numeric_types")]
NumericArray::UInt8(arr) => {
out.extend_from_slice(itoa_buf.format(arr.data.as_ref()[row]).as_bytes());
}
#[cfg(feature = "extended_numeric_types")]
NumericArray::UInt16(arr) => {
out.extend_from_slice(itoa_buf.format(arr.data.as_ref()[row]).as_bytes());
}
NumericArray::UInt32(arr) => {
out.extend_from_slice(itoa_buf.format(arr.data.as_ref()[row]).as_bytes());
}
NumericArray::UInt64(arr) => {
out.extend_from_slice(itoa_buf.format(arr.data.as_ref()[row]).as_bytes());
}
NumericArray::Float32(arr) => {
out.extend_from_slice(ryu_buf.format(arr.data.as_ref()[row]).as_bytes());
}
NumericArray::Float64(arr) => {
out.extend_from_slice(ryu_buf.format(arr.data.as_ref()[row]).as_bytes());
}
_ => {
out.extend_from_slice(b"<unsupported>");
}
},
Array::BooleanArray(arr) => {
out.extend_from_slice(if arr.data.get(row) { b"true" } else { b"false" });
}
Array::TextArray(TextArray::String32(arr)) => {
append_str_cell(
&mut out,
arr.data.as_ref(),
arr.offsets.as_ref(),
row,
delimiter,
quote,
);
}
#[cfg(feature = "large_string")]
Array::TextArray(TextArray::String64(arr)) => {
append_str_cell(
&mut out,
arr.data.as_ref(),
arr.offsets.as_ref(),
row,
delimiter,
quote,
);
}
#[cfg(any(
not(feature = "default_categorical_8"),
feature = "extended_categorical"
))]
Array::TextArray(TextArray::Categorical32(arr)) => {
let idx = arr.data.as_ref()[row] as usize;
let s = arr
.unique_values
.get(idx)
.map(String::as_str)
.unwrap_or("<invalid>");
append_csv_string(&mut out, s.as_bytes(), delimiter, quote);
}
#[cfg(feature = "default_categorical_8")]
Array::TextArray(TextArray::Categorical8(arr)) => {
let idx = arr.data.as_ref()[row] as usize;
let s = arr
.unique_values
.get(idx)
.map(String::as_str)
.unwrap_or("<invalid>");
append_csv_string(&mut out, s.as_bytes(), delimiter, quote);
}
#[cfg(feature = "extended_categorical")]
Array::TextArray(TextArray::Categorical16(arr)) => {
let idx = arr.data.as_ref()[row] as usize;
let s = arr
.unique_values
.get(idx)
.map(String::as_str)
.unwrap_or("<invalid>");
append_csv_string(&mut out, s.as_bytes(), delimiter, quote);
}
#[cfg(feature = "extended_categorical")]
Array::TextArray(TextArray::Categorical64(arr)) => {
let idx = arr.data.as_ref()[row] as usize;
let s = arr
.unique_values
.get(idx)
.map(String::as_str)
.unwrap_or("<invalid>");
append_csv_string(&mut out, s.as_bytes(), delimiter, quote);
}
#[cfg(feature = "datetime")]
Array::TemporalArray(temp) => match temp {
minarrow::TemporalArray::Datetime32(arr) => {
out.extend_from_slice(itoa_buf.format(arr.data.as_ref()[row]).as_bytes());
}
minarrow::TemporalArray::Datetime64(arr) => {
out.extend_from_slice(itoa_buf.format(arr.data.as_ref()[row]).as_bytes());
}
minarrow::TemporalArray::Null => {
out.extend_from_slice(b"<null_temporal>");
}
},
_ => {
out.extend_from_slice(b"<unsupported>");
}
}
}
out.push(b'\n');
}
writer.write_all(&out)
}
pub fn encode_supertable_csv<W: Write>(
supertable: &SuperTable,
mut writer: W,
options: &CsvEncodeOptions,
) -> io::Result<()> {
let mut opts = options.clone();
for (i, batch) in supertable.batches.iter().enumerate() {
opts.write_header = if i == 0 { options.write_header } else { false };
encode_table_csv(batch, &mut writer, &opts)?;
}
Ok(())
}
#[cfg(test)]
mod tests {
use minarrow::{
Array, ArrowType, Bitmask, Buffer, Field, FieldArray, NumericArray, Table, TextArray, vec64,
};
use super::*;
fn make_test_table() -> Table {
let int_col = FieldArray {
field: Field {
name: "ints".to_string(),
dtype: minarrow::ArrowType::Int32,
nullable: true, metadata: Default::default(),
}
.into(),
array: Array::NumericArray(NumericArray::Int32(
minarrow::IntegerArray {
data: Buffer::from(vec64![1, 2, 3, 4]),
null_mask: Some(Bitmask::from_bools(&[true, false, true, true])), }
.into(),
)),
null_count: 1,
};
let str_col = FieldArray {
field: Field {
name: "strings".to_string(),
dtype: minarrow::ArrowType::String,
nullable: true,
metadata: Default::default(),
}
.into(),
array: Array::TextArray(TextArray::String32(
minarrow::StringArray {
offsets: Buffer::from(vec64![0u32, 5, 9, 14, 18]),
data: Buffer::from_vec64("helloabcdworldrust".as_bytes().into()),
null_mask: Some(Bitmask::from_bools(&[true, false, true, true])),
}
.into(),
)),
null_count: 1,
};
Table {
name: "test".to_string(),
cols: vec![int_col, str_col],
n_rows: 4,
..Default::default()
}
}
#[test]
fn test_encode_table_csv_basic() {
let table = make_test_table();
let mut out = Vec::new();
let opts = CsvEncodeOptions::default();
encode_table_csv(&table, &mut out, &opts).unwrap();
let csv = String::from_utf8(out).unwrap();
assert!(csv.contains("ints,strings"));
assert!(csv.contains("hello"));
assert!(csv.contains("\n,\n"));
}
#[test]
fn test_encode_table_csv_custom_delim() {
let table = make_test_table();
let mut out = Vec::new();
let mut opts = CsvEncodeOptions::default();
opts.delimiter = b'\t';
encode_table_csv(&table, &mut out, &opts).unwrap();
let csv = String::from_utf8(out).unwrap();
assert!(csv.contains("\t"));
}
#[test]
fn encode_quotes_field_with_delimiter() {
use minarrow::{Array, Buffer, Field, FieldArray, NumericArray, Table, TextArray, vec64};
use crate::models::encoders::csv::{CsvEncodeOptions, encode_table_csv};
let col1 = FieldArray {
field: Field::new("id", minarrow::ArrowType::Int32, false, None).into(),
array: Array::NumericArray(NumericArray::Int32(
minarrow::IntegerArray {
data: Buffer::from(vec64![1]),
null_mask: None,
}
.into(),
)),
null_count: 0,
};
let col2_str = "needs,quotes"; let col2 = FieldArray {
field: Field::new("txt", minarrow::ArrowType::String, false, None).into(),
array: Array::TextArray(TextArray::String32(
minarrow::StringArray {
offsets: Buffer::from(vec64![0u32, col2_str.len() as u32]),
data: Buffer::from_vec64(col2_str.as_bytes().into()),
null_mask: None,
}
.into(),
)),
null_count: 0,
};
let tbl = Table {
name: "".into(),
cols: vec![col1, col2],
n_rows: 1,
..Default::default()
};
let mut out = Vec::new();
encode_table_csv(&tbl, &mut out, &CsvEncodeOptions::default()).unwrap();
let s = String::from_utf8(out).unwrap();
assert!(s.contains("\"needs,quotes\"")); }
#[test]
fn encode_doubles_embedded_quotes() {
let col = FieldArray {
field: Field::new("txt", minarrow::ArrowType::String, false, None).into(),
array: Array::TextArray(TextArray::String32(
minarrow::StringArray {
offsets: Buffer::from(vec64![0u32, 9]),
data: Buffer::from_vec64(b"he\"llo\",x".to_vec().into()),
null_mask: None,
}
.into(),
)),
null_count: 0,
};
let tbl = Table {
name: "".into(),
cols: vec![col],
n_rows: 1,
..Default::default()
};
let mut out = Vec::new();
encode_table_csv(&tbl, &mut out, &CsvEncodeOptions::default()).unwrap();
let s = String::from_utf8(out).unwrap();
assert!(s.contains("\"he\"\"llo\"\",x\""), "got: {s}");
}
#[test]
fn encode_decode_custom_null() {
use crate::models::decoders::csv::*;
use crate::models::encoders::csv::*;
let mut opts = CsvEncodeOptions::default();
opts.null_repr = "NULL";
use minarrow::{
Array, ArrowType, Bitmask, Field, FieldArray, IntegerArray, NumericArray, Table,
};
use std::sync::Arc;
let field = Field {
name: "int32".to_string(),
dtype: ArrowType::Int32,
nullable: true,
metadata: Default::default(),
};
let null_mask = Bitmask::from_bytes(&[0b00000000], 1); let array = Array::NumericArray(NumericArray::Int32(Arc::new(IntegerArray {
data: Buffer::from(minarrow::Vec64::from_slice(&[42i32])), null_mask: Some(null_mask),
})));
let col = FieldArray::new(field, array);
let tbl = Table {
cols: vec![col],
n_rows: 1,
name: "test_null".to_string(),
..Default::default()
};
let mut buf = Vec::new();
encode_table_csv(&tbl, &mut buf, &opts).unwrap();
let mut dec = CsvDecodeOptions::default();
dec.nulls = vec!["NULL"];
let parsed = decode_csv(std::io::Cursor::new(&buf), &dec).unwrap();
assert_eq!(parsed.cols[0].null_count, 1);
}
#[test]
fn test_csv_decoder_mask_semantics() {
use crate::models::decoders::csv::*;
use minarrow::MaskedArray;
let csv = b"col\nvalid\n\nvalid2\n"; let opts = CsvDecodeOptions::default();
let table = decode_csv(std::io::Cursor::new(csv.as_ref()), &opts).unwrap();
assert_eq!(table.cols[0].null_count, 1);
let Array::TextArray(TextArray::String32(arr)) = &table.cols[0].array else {
panic!("Expected String32 array");
};
let mask = arr.null_mask.as_ref().unwrap();
assert_eq!(mask.len(), 3);
assert_eq!(mask.count_ones(), 2, "2 valid rows");
assert_eq!(mask.count_zeros(), 1, "1 null row");
assert_eq!(arr.null_count(), 1);
assert!(mask.get(0));
assert!(!mask.get(1));
assert!(mask.get(2));
}
#[test]
fn test_null_mask_interpretation_mixed_nulls() {
use minarrow::{
Array, ArrowType, Bitmask, Field, FieldArray, IntegerArray, NumericArray, Table,
};
use std::sync::Arc;
let field = Field {
name: "mixed_nulls".to_string(),
dtype: ArrowType::Int32,
nullable: true,
metadata: Default::default(),
};
let null_mask = Bitmask::from_bytes(&[0b00000101], 4);
let array = Array::NumericArray(NumericArray::Int32(Arc::new(IntegerArray {
data: Buffer::from(minarrow::Vec64::from_slice(&[10i32, 999i32, 30i32, 999i32])),
null_mask: Some(null_mask),
})));
let col = FieldArray::new(field, array);
let tbl = Table {
cols: vec![col],
n_rows: 4,
name: "mixed_null_test".to_string(),
..Default::default()
};
assert_eq!(tbl.cols[0].null_count, 2, "Expected 2 nulls");
let mut opts = CsvEncodeOptions::default();
opts.null_repr = "NULL";
let mut buf = Vec::new();
encode_table_csv(&tbl, &mut buf, &opts).unwrap();
let csv_output = String::from_utf8(buf).unwrap();
assert_eq!(csv_output, "mixed_nulls\n10\nNULL\n30\nNULL\n");
}
#[test]
fn test_null_mask_interpretation_all_nulls() {
use minarrow::{
Array, ArrowType, Bitmask, Field, FieldArray, IntegerArray, NumericArray, Table,
};
use std::sync::Arc;
let field = Field {
name: "all_nulls".to_string(),
dtype: ArrowType::Int32,
nullable: true,
metadata: Default::default(),
};
let null_mask = Bitmask::from_bytes(&[0b00000000], 3);
let array = Array::NumericArray(NumericArray::Int32(Arc::new(IntegerArray {
data: Buffer::from(minarrow::Vec64::from_slice(&[999i32, 999i32, 999i32])),
null_mask: Some(null_mask),
})));
let col = FieldArray::new(field, array);
let tbl = Table {
cols: vec![col],
n_rows: 3,
name: "all_null_test".to_string(),
..Default::default()
};
assert_eq!(tbl.cols[0].null_count, 3, "Expected 3 nulls");
let mut opts = CsvEncodeOptions::default();
opts.null_repr = "NULL";
let mut buf = Vec::new();
encode_table_csv(&tbl, &mut buf, &opts).unwrap();
let csv_output = String::from_utf8(buf).unwrap();
assert_eq!(csv_output, "all_nulls\nNULL\nNULL\nNULL\n");
}
#[test]
fn categorical_roundtrip() {
use crate::models::decoders::csv::*;
use crate::models::encoders::csv::*;
let csv = b"id,fruit\n1,apple\n2,banana\n3,apple\n";
let mut opts = CsvDecodeOptions::default();
opts.categorical_cols.insert("fruit".into());
let tbl = decode_csv(std::io::Cursor::new(csv.as_ref()), &opts).unwrap();
assert!(matches!(tbl.cols[1].field.dtype, ArrowType::Dictionary(_)));
let mut out = Vec::new();
encode_table_csv(&tbl, &mut out, &CsvEncodeOptions::default()).unwrap();
let out_str = String::from_utf8(out).unwrap();
assert!(out_str.contains("apple"));
assert!(out_str.contains("banana"));
}
}