Skip to main content

alopex_sql/executor/bulk/
parquet.rs

1use 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
19/// Parquet リーダー(Arrow 経由でスキーマ抽出とデータ読み込み)。
20pub 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        // TODO: バッチサイズを open 引数で受け取れるようにし、呼び出し側で柔軟に制御できるようにする。
51
52        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}