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 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 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}