Skip to main content

ahri_tre_tabular/
columns.rs

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}