alopex_sql/executor/bulk/
parquet.rs1use std::fs::File;
2
3use arrow_array::types::IntervalMonthDayNanoType;
4use arrow_array::{
5 Array, BinaryArray, BooleanArray, Date32Array, Decimal128Array, Float32Array, Float64Array,
6 Int32Array, Int64Array, IntervalMonthDayNanoArray, LargeBinaryArray, LargeListArray, ListArray,
7 MapArray, StringArray, StructArray, Time64MicrosecondArray, TimestampMicrosecondArray,
8};
9use arrow_schema::{DataType as ArrowDataType, IntervalUnit, TimeUnit};
10use parquet::arrow::arrow_reader::{ParquetRecordBatchReader, ParquetRecordBatchReaderBuilder};
11
12use crate::catalog::TableMetadata;
13use crate::executor::{ExecutorError, Result};
14use crate::planner::types::ResolvedType;
15use crate::storage::SqlValue;
16
17use super::{BulkReader, CopyField, CopySchema};
18
19pub struct ParquetReader {
21 schema: CopySchema,
22 target_types: Vec<ResolvedType>,
23 reader: ParquetRecordBatchReader,
24 buffer: Option<Vec<Vec<SqlValue>>>,
25}
26
27impl ParquetReader {
28 pub fn open(path: &str, table_meta: &TableMetadata, _header: bool) -> Result<Self> {
29 let file = File::open(path)
30 .map_err(|e| ExecutorError::BulkLoad(format!("failed to open parquet: {e}")))?;
31
32 let builder = ParquetRecordBatchReaderBuilder::try_new(file).map_err(|e| {
33 ExecutorError::BulkLoad(format!("failed to read parquet metadata: {e}"))
34 })?;
35
36 let arrow_schema = builder.schema();
37 let mut fields = Vec::with_capacity(arrow_schema.fields().len());
38 for f in arrow_schema.fields() {
39 let ty = map_arrow_type(f.data_type())?;
40 fields.push(CopyField {
41 name: Some(f.name().clone()),
42 data_type: Some(ty),
43 });
44 }
45
46 let reader = builder
47 .with_batch_size(1024)
48 .build()
49 .map_err(|e| ExecutorError::BulkLoad(format!("failed to build parquet reader: {e}")))?;
50 let target_types: Vec<ResolvedType> = table_meta
53 .columns
54 .iter()
55 .map(|c| c.data_type.clone())
56 .collect();
57
58 Ok(Self {
59 schema: CopySchema { fields },
60 target_types,
61 reader,
62 buffer: None,
63 })
64 }
65}
66
67impl BulkReader for ParquetReader {
68 fn schema(&self) -> &CopySchema {
69 &self.schema
70 }
71
72 fn next_batch(&mut self, max_rows: usize) -> Result<Option<Vec<Vec<SqlValue>>>> {
73 let max_rows = max_rows.max(1);
74
75 if let Some(mut buffered) = self.buffer.take() {
76 if buffered.len() > max_rows {
77 let rest = buffered.split_off(max_rows);
78 self.buffer = Some(rest);
79 }
80 return Ok(Some(buffered));
81 }
82
83 let maybe_batch = self.reader.next();
84 let batch = match maybe_batch {
85 Some(b) => b.map_err(|e| {
86 ExecutorError::BulkLoad(format!("failed to read parquet batch: {e}"))
87 })?,
88 None => return Ok(None),
89 };
90
91 let mut rows: Vec<Vec<SqlValue>> = Vec::with_capacity(batch.num_rows());
92 for row_idx in 0..batch.num_rows() {
93 let mut row = Vec::with_capacity(self.schema.fields.len());
94 for col_idx in 0..self.schema.fields.len() {
95 let value = arrow_value_to_sql(
96 batch.column(col_idx).as_ref(),
97 batch.schema().field(col_idx).data_type(),
98 self.target_types
99 .get(col_idx)
100 .ok_or_else(|| ExecutorError::BulkLoad("missing target type".into()))?,
101 row_idx,
102 )?;
103 row.push(value);
104 }
105 rows.push(row);
106 }
107
108 if rows.len() > max_rows {
109 let rest = rows.split_off(max_rows);
110 self.buffer = Some(rest);
111 }
112
113 Ok(Some(rows))
114 }
115}
116
117fn map_arrow_type(dt: &ArrowDataType) -> Result<ResolvedType> {
118 match dt {
119 ArrowDataType::Int32 => Ok(ResolvedType::Integer),
120 ArrowDataType::Int64 => Ok(ResolvedType::BigInt),
121 ArrowDataType::Float32 => Ok(ResolvedType::Float),
122 ArrowDataType::Float64 => Ok(ResolvedType::Double),
123 ArrowDataType::Boolean => Ok(ResolvedType::Boolean),
124 ArrowDataType::Utf8 => Ok(ResolvedType::Text),
125 ArrowDataType::Binary | ArrowDataType::LargeBinary => Ok(ResolvedType::Blob),
126 ArrowDataType::Timestamp(arrow_schema::TimeUnit::Microsecond, _) => {
127 Ok(ResolvedType::Timestamp)
128 }
129 ArrowDataType::Date32 => Ok(ResolvedType::Date),
130 ArrowDataType::Time64(TimeUnit::Microsecond) => Ok(ResolvedType::Time),
131 ArrowDataType::Interval(IntervalUnit::MonthDayNano) => Ok(ResolvedType::Interval),
132 ArrowDataType::Decimal128(precision, scale) if *scale >= 0 => Ok(ResolvedType::Decimal {
133 precision: *precision,
134 scale: *scale as u8,
135 }),
136 other => Err(ExecutorError::BulkLoad(format!(
137 "unsupported parquet/arrow type: {other:?}"
138 ))),
139 }
140}
141
142fn arrow_value_to_sql(
143 array: &dyn Array,
144 dt: &ArrowDataType,
145 expected: &ResolvedType,
146 row_idx: usize,
147) -> Result<SqlValue> {
148 if array.is_null(row_idx) {
149 return Ok(SqlValue::Null);
150 }
151
152 match (dt, expected) {
153 (ArrowDataType::Int32, ResolvedType::Integer) => {
154 let arr = array.as_any().downcast_ref::<Int32Array>().unwrap();
155 Ok(SqlValue::Integer(arr.value(row_idx)))
156 }
157 (ArrowDataType::Int32, ResolvedType::BigInt) => {
158 let arr = array.as_any().downcast_ref::<Int32Array>().unwrap();
159 Ok(SqlValue::BigInt(arr.value(row_idx) as i64))
160 }
161 (ArrowDataType::Int32, ResolvedType::Float) => {
162 let arr = array.as_any().downcast_ref::<Int32Array>().unwrap();
163 Ok(SqlValue::Float(arr.value(row_idx) as f32))
164 }
165 (ArrowDataType::Int32, ResolvedType::Double) => {
166 let arr = array.as_any().downcast_ref::<Int32Array>().unwrap();
167 Ok(SqlValue::Double(arr.value(row_idx) as f64))
168 }
169 (ArrowDataType::Int64, ResolvedType::BigInt) => {
170 let arr = array.as_any().downcast_ref::<Int64Array>().unwrap();
171 Ok(SqlValue::BigInt(arr.value(row_idx)))
172 }
173 (ArrowDataType::Int64, ResolvedType::Double) => {
174 let arr = array.as_any().downcast_ref::<Int64Array>().unwrap();
175 Ok(SqlValue::Double(arr.value(row_idx) as f64))
176 }
177 (ArrowDataType::Float32, ResolvedType::Float) => {
178 let arr = array.as_any().downcast_ref::<Float32Array>().unwrap();
179 Ok(SqlValue::Float(arr.value(row_idx)))
180 }
181 (ArrowDataType::Float32, ResolvedType::Double) => {
182 let arr = array.as_any().downcast_ref::<Float32Array>().unwrap();
183 Ok(SqlValue::Double(arr.value(row_idx) as f64))
184 }
185 (ArrowDataType::Float64, ResolvedType::Double) => {
186 let arr = array.as_any().downcast_ref::<Float64Array>().unwrap();
187 Ok(SqlValue::Double(arr.value(row_idx)))
188 }
189 (ArrowDataType::Boolean, ResolvedType::Boolean) => {
190 let arr = array.as_any().downcast_ref::<BooleanArray>().unwrap();
191 Ok(SqlValue::Boolean(arr.value(row_idx)))
192 }
193 (ArrowDataType::Utf8, ResolvedType::Text) => {
194 let arr = array.as_any().downcast_ref::<StringArray>().unwrap();
195 Ok(SqlValue::Text(arr.value(row_idx).to_string()))
196 }
197 (ArrowDataType::Utf8, ResolvedType::Json) => {
198 let arr = array.as_any().downcast_ref::<StringArray>().unwrap();
199 crate::storage::JsonValue::parse(arr.value(row_idx))
200 .map(SqlValue::Json)
201 .map_err(|error| ExecutorError::BulkLoad(format!("invalid JSON: {error}")))
202 }
203 (
204 ArrowDataType::Utf8,
205 expected
206 @ (ResolvedType::Array(_) | ResolvedType::Map { .. } | ResolvedType::Struct(_)),
207 ) => {
208 let arr = array.as_any().downcast_ref::<StringArray>().unwrap();
209 crate::executor::evaluator::nested::parse_typed_json(arr.value(row_idx), expected)
210 }
211 (ArrowDataType::List(_), ResolvedType::Array(element)) => {
212 let arr = array.as_any().downcast_ref::<ListArray>().unwrap();
213 let values = arr.value(row_idx);
214 (0..values.len())
215 .map(|index| {
216 arrow_value_to_sql(values.as_ref(), values.data_type(), element, index)
217 })
218 .collect::<Result<Vec<_>>>()
219 .map(SqlValue::Array)
220 }
221 (ArrowDataType::LargeList(_), ResolvedType::Array(element)) => {
222 let arr = array.as_any().downcast_ref::<LargeListArray>().unwrap();
223 let values = arr.value(row_idx);
224 (0..values.len())
225 .map(|index| {
226 arrow_value_to_sql(values.as_ref(), values.data_type(), element, index)
227 })
228 .collect::<Result<Vec<_>>>()
229 .map(SqlValue::Array)
230 }
231 (ArrowDataType::Map(_, _), ResolvedType::Map { key, value }) => {
232 let arr = array.as_any().downcast_ref::<MapArray>().unwrap();
233 let entries = arr.value(row_idx);
234 let keys = entries.column(0);
235 let values = entries.column(1);
236 (0..entries.len())
237 .map(|index| {
238 Ok((
239 arrow_value_to_sql(keys.as_ref(), keys.data_type(), key, index)?,
240 arrow_value_to_sql(values.as_ref(), values.data_type(), value, index)?,
241 ))
242 })
243 .collect::<Result<Vec<_>>>()
244 .map(SqlValue::Map)
245 }
246 (ArrowDataType::Struct(arrow_fields), ResolvedType::Struct(expected_fields))
247 if arrow_fields.len() == expected_fields.len() =>
248 {
249 let arr = array.as_any().downcast_ref::<StructArray>().unwrap();
250 expected_fields
251 .iter()
252 .enumerate()
253 .map(|(index, (name, data_type))| {
254 if arrow_fields[index].name() != name {
255 return Err(ExecutorError::BulkLoad(format!(
256 "struct field mismatch: expected {name}, found {}",
257 arrow_fields[index].name()
258 )));
259 }
260 Ok((
261 name.clone(),
262 arrow_value_to_sql(
263 arr.column(index).as_ref(),
264 arrow_fields[index].data_type(),
265 data_type,
266 row_idx,
267 )?,
268 ))
269 })
270 .collect::<Result<Vec<_>>>()
271 .map(SqlValue::Struct)
272 }
273 (ArrowDataType::Binary, ResolvedType::Blob) => {
274 let arr = array.as_any().downcast_ref::<BinaryArray>().unwrap();
275 Ok(SqlValue::Blob(arr.value(row_idx).to_vec()))
276 }
277 (ArrowDataType::LargeBinary, ResolvedType::Blob) => {
278 let arr = array.as_any().downcast_ref::<LargeBinaryArray>().unwrap();
279 Ok(SqlValue::Blob(arr.value(row_idx).to_vec()))
280 }
281 (
282 ArrowDataType::Timestamp(arrow_schema::TimeUnit::Microsecond, _),
283 ResolvedType::Timestamp,
284 ) => {
285 let arr = array
286 .as_any()
287 .downcast_ref::<TimestampMicrosecondArray>()
288 .unwrap();
289 Ok(SqlValue::Timestamp(arr.value(row_idx)))
290 }
291 (ArrowDataType::Date32, ResolvedType::Date) => {
292 let arr = array.as_any().downcast_ref::<Date32Array>().unwrap();
293 Ok(SqlValue::Date(arr.value(row_idx)))
294 }
295 (ArrowDataType::Time64(TimeUnit::Microsecond), ResolvedType::Time) => {
296 let arr = array
297 .as_any()
298 .downcast_ref::<Time64MicrosecondArray>()
299 .unwrap();
300 Ok(SqlValue::Time(arr.value(row_idx)))
301 }
302 (ArrowDataType::Interval(IntervalUnit::MonthDayNano), ResolvedType::Interval) => {
303 let arr = array
304 .as_any()
305 .downcast_ref::<IntervalMonthDayNanoArray>()
306 .unwrap();
307 let value = arr.value(row_idx);
308 let (months, days, nanos) = IntervalMonthDayNanoType::to_parts(value);
309 if nanos % 1_000 != 0 {
310 return Err(ExecutorError::BulkLoad(
311 "interval nanoseconds cannot be represented as microseconds".into(),
312 ));
313 }
314 Ok(SqlValue::Interval {
315 months,
316 days,
317 micros: nanos / 1_000,
318 })
319 }
320 (ArrowDataType::Decimal128(_, arrow_scale), ResolvedType::Decimal { precision, scale })
321 if *arrow_scale >= 0 =>
322 {
323 let arr = array.as_any().downcast_ref::<Decimal128Array>().unwrap();
324 let value = crate::storage::DecimalValue::new(arr.value(row_idx), *arrow_scale as u8)
325 .rescale(*scale)
326 .ok_or_else(|| ExecutorError::BulkLoad("decimal rescale overflow".into()))?;
327 if !value.fits_precision(*precision) {
328 return Err(ExecutorError::BulkLoad("decimal precision overflow".into()));
329 }
330 Ok(SqlValue::Decimal(value))
331 }
332 _ => Err(ExecutorError::BulkLoad(format!(
333 "parquet field type {:?} does not match expected {:?}",
334 dt, expected
335 ))),
336 }
337}
338
339#[cfg(test)]
340mod tests {
341 use super::*;
342 use crate::storage::DecimalValue;
343
344 #[test]
345 fn arrow_list_maps_to_native_array_with_null_elements() {
346 let mut builder =
347 arrow_array::builder::ListBuilder::new(arrow_array::builder::Int32Builder::new());
348 builder.values().append_value(1);
349 builder.values().append_null();
350 builder.values().append_value(2);
351 builder.append(true);
352 let array = builder.finish();
353
354 assert_eq!(
355 arrow_value_to_sql(
356 &array,
357 array.data_type(),
358 &ResolvedType::Array(Box::new(ResolvedType::Integer)),
359 0,
360 )
361 .unwrap(),
362 SqlValue::Array(vec![
363 SqlValue::Integer(1),
364 SqlValue::Null,
365 SqlValue::Integer(2),
366 ])
367 );
368 }
369
370 #[test]
371 fn decimal128_maps_to_exact_sql_decimal() {
372 let array = Decimal128Array::from(vec![Some(12345)])
373 .with_precision_and_scale(10, 3)
374 .unwrap();
375 let ty = ArrowDataType::Decimal128(10, 3);
376 assert_eq!(
377 map_arrow_type(&ty).unwrap(),
378 ResolvedType::Decimal {
379 precision: 10,
380 scale: 3,
381 }
382 );
383 assert_eq!(
384 arrow_value_to_sql(
385 &array,
386 &ty,
387 &ResolvedType::Decimal {
388 precision: 10,
389 scale: 2,
390 },
391 0,
392 )
393 .unwrap(),
394 SqlValue::Decimal(DecimalValue::new(1235, 2))
395 );
396 }
397}