Skip to main content

alopex_embedded/
dataframe_api.rs

1use std::sync::Arc;
2
3use arrow::array::types::IntervalMonthDayNanoType;
4use arrow::array::{
5    ArrayRef, BinaryArray, BooleanArray, Date32Array, Decimal128Array, Float32Array, Float64Array,
6    Int32Array, Int64Array, IntervalMonthDayNanoArray, NullArray, StringArray,
7    Time64MicrosecondArray, TimestampMicrosecondArray,
8};
9use arrow::datatypes::{DataType, Field, IntervalUnit, Schema, TimeUnit};
10use arrow::record_batch::RecordBatch;
11
12use alopex_dataframe::{DataFrame, DataFrameError};
13use alopex_sql::{ColumnInfo, ExecutionResult, QueryResult, ResolvedType, SqlValue};
14
15use crate::{Database, Result, SqlResult, Transaction};
16
17type DfResult<T> = std::result::Result<T, DataFrameError>;
18
19impl Database {
20    /// Execute SQL and return a DataFrame for query results.
21    pub fn query_df(&self, sql: &str) -> Result<DataFrame> {
22        let result = self.execute_sql(sql)?;
23        let df = sql_result_to_dataframe(result)?;
24        Ok(df)
25    }
26}
27
28impl<'a> Transaction<'a> {
29    /// Execute SQL within the transaction and return a DataFrame for query results.
30    pub fn query_df(&mut self, sql: &str) -> Result<DataFrame> {
31        let result = self.execute_sql(sql)?;
32        let df = sql_result_to_dataframe(result)?;
33        Ok(df)
34    }
35}
36
37fn sql_result_to_dataframe(result: SqlResult) -> DfResult<DataFrame> {
38    match result {
39        ExecutionResult::Query(query) => query_result_to_dataframe(query),
40        ExecutionResult::Success | ExecutionResult::RowsAffected(_) => Err(
41            DataFrameError::invalid_operation("query_df requires a SELECT query that returns rows"),
42        ),
43    }
44}
45
46fn query_result_to_dataframe(query: QueryResult) -> DfResult<DataFrame> {
47    let row_count = query.rows.len();
48    let mut fields = Vec::with_capacity(query.columns.len());
49    let mut builders = Vec::with_capacity(query.columns.len());
50
51    for ColumnInfo { name, data_type } in query.columns {
52        let arrow_type = arrow_type_for(&data_type)?;
53        fields.push(Field::new(&name, arrow_type, true));
54        builders.push(ColumnBuilder::new(name, data_type, row_count)?);
55    }
56
57    for row in query.rows {
58        if row.len() != builders.len() {
59            return Err(DataFrameError::schema_mismatch(format!(
60                "row has {} columns, expected {}",
61                row.len(),
62                builders.len()
63            )));
64        }
65
66        for (value, builder) in row.into_iter().zip(builders.iter_mut()) {
67            builder.push(value)?;
68        }
69    }
70
71    let schema = Arc::new(Schema::new(fields));
72    let arrays = builders
73        .into_iter()
74        .map(ColumnBuilder::finish)
75        .collect::<DfResult<Vec<_>>>()?;
76    let batch = RecordBatch::try_new(schema, arrays).map_err(|e| {
77        DataFrameError::schema_mismatch(format!("failed to build RecordBatch: {e}"))
78    })?;
79
80    DataFrame::from_batches(vec![batch])
81}
82
83fn arrow_type_for(ty: &ResolvedType) -> DfResult<DataType> {
84    match ty {
85        ResolvedType::Integer => Ok(DataType::Int32),
86        ResolvedType::BigInt => Ok(DataType::Int64),
87        ResolvedType::Float => Ok(DataType::Float32),
88        ResolvedType::Double => Ok(DataType::Float64),
89        ResolvedType::Text => Ok(DataType::Utf8),
90        ResolvedType::Json => Ok(DataType::Utf8),
91        ResolvedType::Array(_) | ResolvedType::Map { .. } | ResolvedType::Struct(_) => {
92            Ok(DataType::Utf8)
93        }
94        ResolvedType::Blob => Ok(DataType::Binary),
95        ResolvedType::Boolean => Ok(DataType::Boolean),
96        ResolvedType::Timestamp => Ok(DataType::Timestamp(TimeUnit::Microsecond, None)),
97        ResolvedType::Date => Ok(DataType::Date32),
98        ResolvedType::Time => Ok(DataType::Time64(TimeUnit::Microsecond)),
99        ResolvedType::Interval => Ok(DataType::Interval(IntervalUnit::MonthDayNano)),
100        ResolvedType::Decimal { precision, scale } => {
101            Ok(DataType::Decimal128(*precision, *scale as i8))
102        }
103        ResolvedType::Null => Ok(DataType::Null),
104        ResolvedType::Vector { .. } => Err(DataFrameError::invalid_operation(
105            "vector columns are not supported for DataFrame conversion",
106        )),
107    }
108}
109
110struct ColumnBuilder {
111    name: String,
112    expected: ResolvedType,
113    kind: ColumnBuilderKind,
114}
115
116enum ColumnBuilderKind {
117    Int32(Vec<Option<i32>>),
118    Int64(Vec<Option<i64>>),
119    Float32(Vec<Option<f32>>),
120    Float64(Vec<Option<f64>>),
121    Utf8(Vec<Option<String>>),
122    Binary(Vec<Option<Vec<u8>>>),
123    Boolean(Vec<Option<bool>>),
124    Timestamp(Vec<Option<i64>>),
125    Date(Vec<Option<i32>>),
126    Time(Vec<Option<i64>>),
127    Interval(Vec<Option<<IntervalMonthDayNanoType as arrow::array::ArrowPrimitiveType>::Native>>),
128    Decimal(Vec<Option<i128>>, u8, i8),
129    Null(usize),
130}
131
132impl ColumnBuilder {
133    fn new(name: String, expected: ResolvedType, row_count: usize) -> DfResult<Self> {
134        let kind = match expected {
135            ResolvedType::Integer => ColumnBuilderKind::Int32(Vec::with_capacity(row_count)),
136            ResolvedType::BigInt => ColumnBuilderKind::Int64(Vec::with_capacity(row_count)),
137            ResolvedType::Float => ColumnBuilderKind::Float32(Vec::with_capacity(row_count)),
138            ResolvedType::Double => ColumnBuilderKind::Float64(Vec::with_capacity(row_count)),
139            ResolvedType::Text => ColumnBuilderKind::Utf8(Vec::with_capacity(row_count)),
140            ResolvedType::Json => ColumnBuilderKind::Utf8(Vec::with_capacity(row_count)),
141            ResolvedType::Array(_) | ResolvedType::Map { .. } | ResolvedType::Struct(_) => {
142                ColumnBuilderKind::Utf8(Vec::with_capacity(row_count))
143            }
144            ResolvedType::Blob => ColumnBuilderKind::Binary(Vec::with_capacity(row_count)),
145            ResolvedType::Boolean => ColumnBuilderKind::Boolean(Vec::with_capacity(row_count)),
146            ResolvedType::Timestamp => ColumnBuilderKind::Timestamp(Vec::with_capacity(row_count)),
147            ResolvedType::Date => ColumnBuilderKind::Date(Vec::with_capacity(row_count)),
148            ResolvedType::Time => ColumnBuilderKind::Time(Vec::with_capacity(row_count)),
149            ResolvedType::Interval => ColumnBuilderKind::Interval(Vec::with_capacity(row_count)),
150            ResolvedType::Decimal { precision, scale } => {
151                ColumnBuilderKind::Decimal(Vec::with_capacity(row_count), precision, scale as i8)
152            }
153            ResolvedType::Null => ColumnBuilderKind::Null(0),
154            ResolvedType::Vector { .. } => {
155                return Err(DataFrameError::invalid_operation(
156                    "vector columns are not supported for DataFrame conversion",
157                ))
158            }
159        };
160
161        Ok(Self {
162            name,
163            expected,
164            kind,
165        })
166    }
167
168    fn push(&mut self, value: SqlValue) -> DfResult<()> {
169        match (&mut self.kind, value) {
170            (ColumnBuilderKind::Int32(values), SqlValue::Integer(v)) => {
171                values.push(Some(v));
172                Ok(())
173            }
174            (ColumnBuilderKind::Int64(values), SqlValue::BigInt(v)) => {
175                values.push(Some(v));
176                Ok(())
177            }
178            (ColumnBuilderKind::Float32(values), SqlValue::Float(v)) => {
179                values.push(Some(v));
180                Ok(())
181            }
182            (ColumnBuilderKind::Float64(values), SqlValue::Double(v)) => {
183                values.push(Some(v));
184                Ok(())
185            }
186            (ColumnBuilderKind::Utf8(values), SqlValue::Text(v)) => {
187                values.push(Some(v));
188                Ok(())
189            }
190            (ColumnBuilderKind::Utf8(values), SqlValue::Json(v)) => {
191                values.push(Some(v.to_string()));
192                Ok(())
193            }
194            (
195                ColumnBuilderKind::Utf8(values),
196                value @ (SqlValue::Array(_) | SqlValue::Map(_) | SqlValue::Struct(_)),
197            ) => {
198                values.push(Some(value.nested_json_text().ok_or_else(|| {
199                    DataFrameError::invalid_operation("nested value cannot be mapped to Arrow UTF8")
200                })?));
201                Ok(())
202            }
203            (ColumnBuilderKind::Binary(values), SqlValue::Blob(v)) => {
204                values.push(Some(v));
205                Ok(())
206            }
207            (ColumnBuilderKind::Boolean(values), SqlValue::Boolean(v)) => {
208                values.push(Some(v));
209                Ok(())
210            }
211            (ColumnBuilderKind::Timestamp(values), SqlValue::Timestamp(v)) => {
212                values.push(Some(v));
213                Ok(())
214            }
215            (ColumnBuilderKind::Date(values), SqlValue::Date(v)) => {
216                values.push(Some(v));
217                Ok(())
218            }
219            (ColumnBuilderKind::Time(values), SqlValue::Time(v)) => {
220                values.push(Some(v));
221                Ok(())
222            }
223            (
224                ColumnBuilderKind::Interval(values),
225                SqlValue::Interval {
226                    months,
227                    days,
228                    micros,
229                },
230            ) => {
231                let nanos = micros.checked_mul(1_000).ok_or_else(|| {
232                    DataFrameError::invalid_operation("interval nanoseconds overflow Arrow i64")
233                })?;
234                values.push(Some(IntervalMonthDayNanoType::make_value(
235                    months, days, nanos,
236                )));
237                Ok(())
238            }
239            (ColumnBuilderKind::Decimal(values, _, _), SqlValue::Decimal(v)) => {
240                values.push(Some(v.coefficient));
241                Ok(())
242            }
243            (ColumnBuilderKind::Int32(values), SqlValue::Null) => {
244                values.push(None);
245                Ok(())
246            }
247            (ColumnBuilderKind::Int64(values), SqlValue::Null) => {
248                values.push(None);
249                Ok(())
250            }
251            (ColumnBuilderKind::Float32(values), SqlValue::Null) => {
252                values.push(None);
253                Ok(())
254            }
255            (ColumnBuilderKind::Float64(values), SqlValue::Null) => {
256                values.push(None);
257                Ok(())
258            }
259            (ColumnBuilderKind::Utf8(values), SqlValue::Null) => {
260                values.push(None);
261                Ok(())
262            }
263            (ColumnBuilderKind::Binary(values), SqlValue::Null) => {
264                values.push(None);
265                Ok(())
266            }
267            (ColumnBuilderKind::Boolean(values), SqlValue::Null) => {
268                values.push(None);
269                Ok(())
270            }
271            (ColumnBuilderKind::Timestamp(values), SqlValue::Null) => {
272                values.push(None);
273                Ok(())
274            }
275            (ColumnBuilderKind::Date(values), SqlValue::Null) => {
276                values.push(None);
277                Ok(())
278            }
279            (ColumnBuilderKind::Time(values), SqlValue::Null) => {
280                values.push(None);
281                Ok(())
282            }
283            (ColumnBuilderKind::Interval(values), SqlValue::Null) => {
284                values.push(None);
285                Ok(())
286            }
287            (ColumnBuilderKind::Decimal(values, _, _), SqlValue::Null) => {
288                values.push(None);
289                Ok(())
290            }
291            (ColumnBuilderKind::Null(count), SqlValue::Null) => {
292                *count += 1;
293                Ok(())
294            }
295            (_, other) => Err(DataFrameError::type_mismatch(
296                Some(self.name.clone()),
297                self.expected.to_string(),
298                other.type_name().to_string(),
299            )),
300        }
301    }
302
303    fn finish(self) -> DfResult<ArrayRef> {
304        let array: ArrayRef = match self.kind {
305            ColumnBuilderKind::Int32(values) => Arc::new(Int32Array::from(values)),
306            ColumnBuilderKind::Int64(values) => Arc::new(Int64Array::from(values)),
307            ColumnBuilderKind::Float32(values) => Arc::new(Float32Array::from(values)),
308            ColumnBuilderKind::Float64(values) => Arc::new(Float64Array::from(values)),
309            ColumnBuilderKind::Utf8(values) => Arc::new(StringArray::from(values)),
310            ColumnBuilderKind::Binary(values) => {
311                let slices: Vec<Option<&[u8]>> = values.iter().map(|v| v.as_deref()).collect();
312                Arc::new(BinaryArray::from(slices))
313            }
314            ColumnBuilderKind::Boolean(values) => Arc::new(BooleanArray::from(values)),
315            ColumnBuilderKind::Timestamp(values) => {
316                Arc::new(TimestampMicrosecondArray::from(values))
317            }
318            ColumnBuilderKind::Date(values) => Arc::new(Date32Array::from(values)),
319            ColumnBuilderKind::Time(values) => Arc::new(Time64MicrosecondArray::from(values)),
320            ColumnBuilderKind::Interval(values) => {
321                Arc::new(IntervalMonthDayNanoArray::from(values))
322            }
323            ColumnBuilderKind::Decimal(values, precision, scale) => Arc::new(
324                Decimal128Array::from(values)
325                    .with_precision_and_scale(precision, scale)
326                    .map_err(|error| {
327                        DataFrameError::schema_mismatch(format!(
328                            "invalid DECIMAL({precision},{scale}) Arrow type: {error}"
329                        ))
330                    })?,
331            ),
332            ColumnBuilderKind::Null(len) => Arc::new(NullArray::new(len)),
333        };
334
335        Ok(array)
336    }
337}