1use std::sync::Arc;
2
3use arrow_array::{
4 Array, ArrayRef, BooleanArray, Float64Array, Int64Array, RecordBatch, StringArray,
5};
6use arrow_schema::{DataType, Field, Schema};
7
8use crate::{Result, TabularError};
9
10#[derive(Debug, Clone, PartialEq)]
11pub enum TabularColumn {
12 Utf8 {
13 name: String,
14 values: Vec<Option<String>>,
15 },
16 Int64 {
17 name: String,
18 values: Vec<Option<i64>>,
19 },
20 Float64 {
21 name: String,
22 values: Vec<Option<f64>>,
23 },
24 Bool {
25 name: String,
26 values: Vec<Option<bool>>,
27 },
28}
29
30impl TabularColumn {
31 pub fn name(&self) -> &str {
32 match self {
33 Self::Utf8 { name, .. }
34 | Self::Int64 { name, .. }
35 | Self::Float64 { name, .. }
36 | Self::Bool { name, .. } => name,
37 }
38 }
39
40 pub fn len(&self) -> usize {
41 match self {
42 Self::Utf8 { values, .. } => values.len(),
43 Self::Int64 { values, .. } => values.len(),
44 Self::Float64 { values, .. } => values.len(),
45 Self::Bool { values, .. } => values.len(),
46 }
47 }
48
49 pub fn is_empty(&self) -> bool {
50 self.len() == 0
51 }
52
53 fn nullable(&self) -> bool {
54 match self {
55 Self::Utf8 { values, .. } => values.iter().any(Option::is_none),
56 Self::Int64 { values, .. } => values.iter().any(Option::is_none),
57 Self::Float64 { values, .. } => values.iter().any(Option::is_none),
58 Self::Bool { values, .. } => values.iter().any(Option::is_none),
59 }
60 }
61
62 fn field(&self) -> Field {
63 let data_type = match self {
64 Self::Utf8 { .. } => DataType::Utf8,
65 Self::Int64 { .. } => DataType::Int64,
66 Self::Float64 { .. } => DataType::Float64,
67 Self::Bool { .. } => DataType::Boolean,
68 };
69 Field::new(self.name(), data_type, self.nullable())
70 }
71
72 fn array(&self) -> ArrayRef {
73 match self {
74 Self::Utf8 { values, .. } => Arc::new(StringArray::from_iter(
75 values.iter().map(|value| value.as_deref()),
76 )),
77 Self::Int64 { values, .. } => Arc::new(Int64Array::from(values.clone())),
78 Self::Float64 { values, .. } => Arc::new(Float64Array::from(values.clone())),
79 Self::Bool { values, .. } => Arc::new(BooleanArray::from(values.clone())),
80 }
81 }
82}
83
84pub fn columns_to_record_batch(columns: Vec<TabularColumn>) -> Result<RecordBatch> {
85 let first = columns.first().ok_or(TabularError::EmptyColumns)?;
86 let expected = first.len();
87 for column in &columns {
88 let actual = column.len();
89 if actual != expected {
90 return Err(TabularError::ColumnLengthMismatch {
91 name: column.name().to_string(),
92 expected,
93 actual,
94 });
95 }
96 }
97
98 let fields: Vec<Field> = columns.iter().map(TabularColumn::field).collect();
99 let arrays: Vec<ArrayRef> = columns.iter().map(TabularColumn::array).collect();
100 Ok(RecordBatch::try_new(Arc::new(Schema::new(fields)), arrays)?)
101}
102
103pub fn record_batch_to_columns(batch: &RecordBatch) -> Result<Vec<TabularColumn>> {
104 batch
105 .schema()
106 .fields()
107 .iter()
108 .zip(batch.columns())
109 .map(|(field, array)| array_to_column(field, array))
110 .collect()
111}
112
113fn array_to_column(field: &Field, array: &ArrayRef) -> Result<TabularColumn> {
114 match field.data_type() {
115 DataType::Utf8 => {
116 let values = downcast::<StringArray>(field, array)?
117 .iter()
118 .map(|value| value.map(ToString::to_string))
119 .collect();
120 Ok(TabularColumn::Utf8 {
121 name: field.name().to_string(),
122 values,
123 })
124 }
125 DataType::Int64 => {
126 let values = downcast::<Int64Array>(field, array)?.iter().collect();
127 Ok(TabularColumn::Int64 {
128 name: field.name().to_string(),
129 values,
130 })
131 }
132 DataType::Float64 => {
133 let values = downcast::<Float64Array>(field, array)?.iter().collect();
134 Ok(TabularColumn::Float64 {
135 name: field.name().to_string(),
136 values,
137 })
138 }
139 DataType::Boolean => {
140 let values = downcast::<BooleanArray>(field, array)?.iter().collect();
141 Ok(TabularColumn::Bool {
142 name: field.name().to_string(),
143 values,
144 })
145 }
146 data_type => Err(TabularError::UnsupportedDataType {
147 field: field.name().to_string(),
148 data_type: data_type.clone(),
149 }),
150 }
151}
152
153fn downcast<'a, T: Array + 'static>(field: &Field, array: &'a ArrayRef) -> Result<&'a T> {
154 array
155 .as_any()
156 .downcast_ref::<T>()
157 .ok_or_else(|| TabularError::ColumnDowncast {
158 field: field.name().to_string(),
159 data_type: field.data_type().clone(),
160 })
161}