Skip to main content

arrow_pg/
list_encoder.rs

1use std::{str::FromStr, sync::Arc};
2
3#[cfg(not(feature = "datafusion"))]
4use arrow::{
5    array::{
6        Array, BinaryArray, BinaryViewArray, BooleanArray, Date32Array, Date64Array,
7        Decimal128Array, Decimal256Array, DurationMicrosecondArray, DurationMillisecondArray,
8        DurationNanosecondArray, DurationSecondArray, IntervalDayTimeArray,
9        IntervalMonthDayNanoArray, IntervalYearMonthArray, LargeBinaryArray, LargeListArray,
10        LargeStringArray, ListArray, MapArray, PrimitiveArray, StringArray, StringViewArray,
11        Time32MillisecondArray, Time32SecondArray, Time64MicrosecondArray, Time64NanosecondArray,
12        TimestampMicrosecondArray, TimestampMillisecondArray, TimestampNanosecondArray,
13        TimestampSecondArray, timezone::Tz,
14    },
15    datatypes::{
16        DataType, Date32Type, Date64Type, Float32Type, Float64Type, Int8Type, Int16Type, Int32Type,
17        Int64Type, IntervalDayTimeType, IntervalMonthDayNanoType, IntervalUnit,
18        Time32MillisecondType, Time32SecondType, Time64MicrosecondType, Time64NanosecondType,
19        TimeUnit, UInt8Type, UInt16Type, UInt32Type, UInt64Type,
20    },
21    temporal_conversions::{as_date, as_time},
22};
23#[cfg(feature = "datafusion")]
24use datafusion::arrow::{
25    array::{
26        Array, BinaryArray, BinaryViewArray, BooleanArray, Date32Array, Date64Array,
27        Decimal128Array, Decimal256Array, DurationMicrosecondArray, DurationMillisecondArray,
28        DurationNanosecondArray, DurationSecondArray, IntervalDayTimeArray,
29        IntervalMonthDayNanoArray, IntervalYearMonthArray, LargeBinaryArray, LargeListArray,
30        LargeStringArray, ListArray, MapArray, PrimitiveArray, StringArray, StringViewArray,
31        Time32MillisecondArray, Time32SecondArray, Time64MicrosecondArray, Time64NanosecondArray,
32        TimestampMicrosecondArray, TimestampMillisecondArray, TimestampNanosecondArray,
33        TimestampSecondArray, timezone::Tz,
34    },
35    datatypes::{
36        DataType, Date32Type, Date64Type, Float32Type, Float64Type, Int8Type, Int16Type, Int32Type,
37        Int64Type, IntervalDayTimeType, IntervalMonthDayNanoType, IntervalUnit,
38        Time32MillisecondType, Time32SecondType, Time64MicrosecondType, Time64NanosecondType,
39        TimeUnit, UInt8Type, UInt16Type, UInt32Type, UInt64Type,
40    },
41    temporal_conversions::{as_date, as_time},
42};
43
44use chrono::{DateTime, TimeZone, Utc};
45use pg_interval::Interval as PgInterval;
46use pgwire::api::results::FieldInfo;
47use pgwire::error::{PgWireError, PgWireResult};
48use rust_decimal::Decimal;
49
50use crate::encoder::Encoder;
51use crate::error::ToSqlError;
52use crate::struct_encoder::encode_structs;
53
54fn get_bool_list_value(arr: &Arc<dyn Array>) -> Vec<Option<bool>> {
55    arr.as_any()
56        .downcast_ref::<BooleanArray>()
57        .unwrap()
58        .iter()
59        .collect()
60}
61
62macro_rules! get_primitive_list_value {
63    ($name:ident, $t:ty, $pt:ty) => {
64        fn $name(arr: &Arc<dyn Array>) -> Vec<Option<$pt>> {
65            arr.as_any()
66                .downcast_ref::<PrimitiveArray<$t>>()
67                .unwrap()
68                .iter()
69                .collect()
70        }
71    };
72
73    ($name:ident, $t:ty, $pt:ty, $f:expr) => {
74        fn $name(arr: &Arc<dyn Array>) -> Vec<Option<$pt>> {
75            arr.as_any()
76                .downcast_ref::<PrimitiveArray<$t>>()
77                .unwrap()
78                .iter()
79                .map(|val| val.map($f))
80                .collect()
81        }
82    };
83}
84
85get_primitive_list_value!(get_i8_list_value, Int8Type, i8);
86get_primitive_list_value!(get_i16_list_value, Int16Type, i16);
87get_primitive_list_value!(get_i32_list_value, Int32Type, i32);
88get_primitive_list_value!(get_i64_list_value, Int64Type, i64);
89get_primitive_list_value!(get_u8_list_value, UInt8Type, i16, |val: u8| { val as i16 });
90get_primitive_list_value!(get_u16_list_value, UInt16Type, i32, |val: u16| {
91    val as i32
92});
93get_primitive_list_value!(get_u32_list_value, UInt32Type, i64, |val: u32| {
94    val as i64
95});
96get_primitive_list_value!(get_u64_list_value, UInt64Type, Decimal, |val: u64| {
97    Decimal::from(val)
98});
99get_primitive_list_value!(get_f32_list_value, Float32Type, f32);
100get_primitive_list_value!(get_f64_list_value, Float64Type, f64);
101
102pub fn encode_list<T: Encoder>(
103    encoder: &mut T,
104    arr: Arc<dyn Array>,
105    pg_field: &FieldInfo,
106) -> PgWireResult<()> {
107    match arr.data_type() {
108        DataType::Null => {
109            // The element type is unknown (e.g. `ARRAY[NULL]`). Align with
110            // PostgreSQL and treat the list as `text[]`, preserving the
111            // number of (null) elements so that `ARRAY[NULL]` yields `{NULL}`
112            // rather than a SQL NULL.
113            let value: Vec<Option<&str>> = (0..arr.len()).map(|_| None).collect();
114            encoder.encode_field(&value, pg_field)?;
115            Ok(())
116        }
117        DataType::Boolean => {
118            encoder.encode_field(&get_bool_list_value(&arr), pg_field)?;
119            Ok(())
120        }
121        DataType::Int8 => {
122            encoder.encode_field(&get_i8_list_value(&arr), pg_field)?;
123            Ok(())
124        }
125        DataType::Int16 => {
126            encoder.encode_field(&get_i16_list_value(&arr), pg_field)?;
127            Ok(())
128        }
129        DataType::Int32 => {
130            encoder.encode_field(&get_i32_list_value(&arr), pg_field)?;
131            Ok(())
132        }
133        DataType::Int64 => {
134            encoder.encode_field(&get_i64_list_value(&arr), pg_field)?;
135            Ok(())
136        }
137        DataType::UInt8 => {
138            encoder.encode_field(&get_u8_list_value(&arr), pg_field)?;
139            Ok(())
140        }
141        DataType::UInt16 => {
142            encoder.encode_field(&get_u16_list_value(&arr), pg_field)?;
143            Ok(())
144        }
145        DataType::UInt32 => {
146            encoder.encode_field(&get_u32_list_value(&arr), pg_field)?;
147            Ok(())
148        }
149        DataType::UInt64 => {
150            encoder.encode_field(&get_u64_list_value(&arr), pg_field)?;
151            Ok(())
152        }
153        DataType::Float32 => {
154            encoder.encode_field(&get_f32_list_value(&arr), pg_field)?;
155            Ok(())
156        }
157        DataType::Float64 => {
158            encoder.encode_field(&get_f64_list_value(&arr), pg_field)?;
159            Ok(())
160        }
161        DataType::Decimal128(_, s) => {
162            let value: Vec<_> = arr
163                .as_any()
164                .downcast_ref::<Decimal128Array>()
165                .unwrap()
166                .iter()
167                .map(|ov| ov.map(|v| Decimal::from_i128_with_scale(v, *s as u32)))
168                .collect();
169            encoder.encode_field(&value, pg_field)
170        }
171        DataType::Utf8 => {
172            let value: Vec<Option<&str>> = arr
173                .as_any()
174                .downcast_ref::<StringArray>()
175                .unwrap()
176                .iter()
177                .collect();
178            encoder.encode_field(&value, pg_field)
179        }
180        DataType::Utf8View => {
181            let value: Vec<Option<&str>> = arr
182                .as_any()
183                .downcast_ref::<StringViewArray>()
184                .unwrap()
185                .iter()
186                .collect();
187            encoder.encode_field(&value, pg_field)
188        }
189        DataType::Binary => {
190            let value: Vec<Option<_>> = arr
191                .as_any()
192                .downcast_ref::<BinaryArray>()
193                .unwrap()
194                .iter()
195                .collect();
196            encoder.encode_field(&value, pg_field)
197        }
198        DataType::LargeBinary => {
199            let value: Vec<Option<_>> = arr
200                .as_any()
201                .downcast_ref::<LargeBinaryArray>()
202                .unwrap()
203                .iter()
204                .collect();
205            encoder.encode_field(&value, pg_field)
206        }
207        DataType::BinaryView => {
208            let value: Vec<Option<_>> = arr
209                .as_any()
210                .downcast_ref::<BinaryViewArray>()
211                .unwrap()
212                .iter()
213                .collect();
214            encoder.encode_field(&value, pg_field)
215        }
216
217        DataType::Date32 => {
218            let value: Vec<Option<_>> = arr
219                .as_any()
220                .downcast_ref::<Date32Array>()
221                .unwrap()
222                .iter()
223                .map(|val| val.and_then(|x| as_date::<Date32Type>(x as i64)))
224                .collect();
225            encoder.encode_field(&value, pg_field)
226        }
227        DataType::Date64 => {
228            let value: Vec<Option<_>> = arr
229                .as_any()
230                .downcast_ref::<Date64Array>()
231                .unwrap()
232                .iter()
233                .map(|val| val.and_then(as_date::<Date64Type>))
234                .collect();
235            encoder.encode_field(&value, pg_field)
236        }
237        DataType::Time32(unit) => match unit {
238            TimeUnit::Second => {
239                let value: Vec<Option<_>> = arr
240                    .as_any()
241                    .downcast_ref::<Time32SecondArray>()
242                    .unwrap()
243                    .iter()
244                    .map(|val| val.and_then(|x| as_time::<Time32SecondType>(x as i64)))
245                    .collect();
246                encoder.encode_field(&value, pg_field)
247            }
248            TimeUnit::Millisecond => {
249                let value: Vec<Option<_>> = arr
250                    .as_any()
251                    .downcast_ref::<Time32MillisecondArray>()
252                    .unwrap()
253                    .iter()
254                    .map(|val| val.and_then(|x| as_time::<Time32MillisecondType>(x as i64)))
255                    .collect();
256                encoder.encode_field(&value, pg_field)
257            }
258            _ => {
259                // Time32 only supports Second and Millisecond in Arrow
260                // Other units are not available, so return an error
261                Err(PgWireError::ApiError("Unsupported Time32 unit".into()))
262            }
263        },
264        DataType::Time64(unit) => match unit {
265            TimeUnit::Microsecond => {
266                let value: Vec<Option<_>> = arr
267                    .as_any()
268                    .downcast_ref::<Time64MicrosecondArray>()
269                    .unwrap()
270                    .iter()
271                    .map(|val| val.and_then(as_time::<Time64MicrosecondType>))
272                    .collect();
273                encoder.encode_field(&value, pg_field)
274            }
275            TimeUnit::Nanosecond => {
276                let value: Vec<Option<_>> = arr
277                    .as_any()
278                    .downcast_ref::<Time64NanosecondArray>()
279                    .unwrap()
280                    .iter()
281                    .map(|val| val.and_then(as_time::<Time64NanosecondType>))
282                    .collect();
283                encoder.encode_field(&value, pg_field)
284            }
285            _ => {
286                // Time64 only supports Microsecond and Nanosecond in Arrow
287                // Other units are not available, so return an error
288                Err(PgWireError::ApiError("Unsupported Time64 unit".into()))
289            }
290        },
291        DataType::Timestamp(unit, timezone) => match unit {
292            TimeUnit::Second => {
293                let array_iter = arr
294                    .as_any()
295                    .downcast_ref::<TimestampSecondArray>()
296                    .unwrap()
297                    .iter();
298
299                if let Some(tz) = timezone {
300                    let tz = Tz::from_str(tz.as_ref())
301                        .map_err(|e| PgWireError::ApiError(ToSqlError::from(e)))?;
302                    let value: Vec<_> = array_iter
303                        .map(|i| {
304                            i.and_then(|i| {
305                                DateTime::from_timestamp(i, 0).map(|dt| {
306                                    Utc.from_utc_datetime(&dt.naive_utc())
307                                        .with_timezone(&tz)
308                                        .fixed_offset()
309                                })
310                            })
311                        })
312                        .collect();
313                    encoder.encode_field(&value, pg_field)
314                } else {
315                    let value: Vec<_> = array_iter
316                        .map(|i| {
317                            i.and_then(|i| DateTime::from_timestamp(i, 0).map(|dt| dt.naive_utc()))
318                        })
319                        .collect();
320                    encoder.encode_field(&value, pg_field)
321                }
322            }
323            TimeUnit::Millisecond => {
324                let array_iter = arr
325                    .as_any()
326                    .downcast_ref::<TimestampMillisecondArray>()
327                    .unwrap()
328                    .iter();
329
330                if let Some(tz) = timezone {
331                    let tz = Tz::from_str(tz.as_ref()).map_err(ToSqlError::from)?;
332                    let value: Vec<_> = array_iter
333                        .map(|i| {
334                            i.and_then(|i| {
335                                DateTime::from_timestamp_millis(i).map(|dt| {
336                                    Utc.from_utc_datetime(&dt.naive_utc())
337                                        .with_timezone(&tz)
338                                        .fixed_offset()
339                                })
340                            })
341                        })
342                        .collect();
343                    encoder.encode_field(&value, pg_field)
344                } else {
345                    let value: Vec<_> = array_iter
346                        .map(|i| {
347                            i.and_then(|i| {
348                                DateTime::from_timestamp_millis(i).map(|dt| dt.naive_utc())
349                            })
350                        })
351                        .collect();
352                    encoder.encode_field(&value, pg_field)
353                }
354            }
355            TimeUnit::Microsecond => {
356                let array_iter = arr
357                    .as_any()
358                    .downcast_ref::<TimestampMicrosecondArray>()
359                    .unwrap()
360                    .iter();
361
362                if let Some(tz) = timezone {
363                    let tz = Tz::from_str(tz.as_ref()).map_err(ToSqlError::from)?;
364                    let value: Vec<_> = array_iter
365                        .map(|i| {
366                            i.and_then(|i| {
367                                DateTime::from_timestamp_micros(i).map(|dt| {
368                                    Utc.from_utc_datetime(&dt.naive_utc())
369                                        .with_timezone(&tz)
370                                        .fixed_offset()
371                                })
372                            })
373                        })
374                        .collect();
375                    encoder.encode_field(&value, pg_field)
376                } else {
377                    let value: Vec<_> = array_iter
378                        .map(|i| {
379                            i.and_then(|i| {
380                                DateTime::from_timestamp_micros(i).map(|dt| dt.naive_utc())
381                            })
382                        })
383                        .collect();
384                    encoder.encode_field(&value, pg_field)
385                }
386            }
387            TimeUnit::Nanosecond => {
388                let array_iter = arr
389                    .as_any()
390                    .downcast_ref::<TimestampNanosecondArray>()
391                    .unwrap()
392                    .iter();
393
394                if let Some(tz) = timezone {
395                    let tz = Tz::from_str(tz.as_ref()).map_err(ToSqlError::from)?;
396                    let value: Vec<_> = array_iter
397                        .map(|i| {
398                            i.map(|i| {
399                                Utc.from_utc_datetime(
400                                    &DateTime::from_timestamp_nanos(i).naive_utc(),
401                                )
402                                .with_timezone(&tz)
403                                .fixed_offset()
404                            })
405                        })
406                        .collect();
407                    encoder.encode_field(&value, pg_field)
408                } else {
409                    let value: Vec<_> = array_iter
410                        .map(|i| i.map(|i| DateTime::from_timestamp_nanos(i).naive_utc()))
411                        .collect();
412                    encoder.encode_field(&value, pg_field)
413                }
414            }
415        },
416        DataType::Struct(arrow_fields) => encode_structs(encoder, &arr, arrow_fields, pg_field),
417        DataType::LargeUtf8 => {
418            let value: Vec<Option<&str>> = arr
419                .as_any()
420                .downcast_ref::<LargeStringArray>()
421                .unwrap()
422                .iter()
423                .collect();
424            encoder.encode_field(&value, pg_field)?;
425            Ok(())
426        }
427        DataType::Decimal256(_, s) => {
428            // Convert Decimal256 to string representation for now
429            // since rust_decimal doesn't support 256-bit decimals
430            let decimal_array = arr.as_any().downcast_ref::<Decimal256Array>().unwrap();
431            let value: Vec<Option<String>> = (0..decimal_array.len())
432                .map(|i| {
433                    if decimal_array.is_null(i) {
434                        None
435                    } else {
436                        // Convert to string representation
437                        let raw_value = decimal_array.value(i);
438                        let scale = *s as u32;
439                        // Convert i256 to string and handle decimal placement manually
440                        let value_str = raw_value.to_string();
441                        if scale == 0 {
442                            Some(value_str)
443                        } else {
444                            // Insert decimal point
445                            let mut chars: Vec<char> = value_str.chars().collect();
446                            if chars.len() <= scale as usize {
447                                // Prepend zeros if needed
448                                let zeros_needed = scale as usize - chars.len() + 1;
449                                chars.splice(0..0, std::iter::repeat_n('0', zeros_needed));
450                                chars.insert(1, '.');
451                            } else {
452                                let decimal_pos = chars.len() - scale as usize;
453                                chars.insert(decimal_pos, '.');
454                            }
455                            Some(chars.into_iter().collect())
456                        }
457                    }
458                })
459                .collect();
460            encoder.encode_field(&value, pg_field)?;
461            Ok(())
462        }
463        DataType::Duration(unit) => match unit {
464            TimeUnit::Second => {
465                let value: Vec<Option<PgInterval>> = arr
466                    .as_any()
467                    .downcast_ref::<DurationSecondArray>()
468                    .unwrap()
469                    .iter()
470                    .map(|val| val.map(|v| PgInterval::new(0, 0, v * 1_000_000i64)))
471                    .collect();
472                encoder.encode_field(&value, pg_field)?;
473                Ok(())
474            }
475            TimeUnit::Millisecond => {
476                let value: Vec<Option<PgInterval>> = arr
477                    .as_any()
478                    .downcast_ref::<DurationMillisecondArray>()
479                    .unwrap()
480                    .iter()
481                    .map(|val| val.map(|v| PgInterval::new(0, 0, v * 1_000i64)))
482                    .collect();
483                encoder.encode_field(&value, pg_field)?;
484                Ok(())
485            }
486            TimeUnit::Microsecond => {
487                let value: Vec<Option<PgInterval>> = arr
488                    .as_any()
489                    .downcast_ref::<DurationMicrosecondArray>()
490                    .unwrap()
491                    .iter()
492                    .map(|val| val.map(|v| PgInterval::new(0, 0, v)))
493                    .collect();
494                encoder.encode_field(&value, pg_field)?;
495                Ok(())
496            }
497            TimeUnit::Nanosecond => {
498                let value: Vec<Option<PgInterval>> = arr
499                    .as_any()
500                    .downcast_ref::<DurationNanosecondArray>()
501                    .unwrap()
502                    .iter()
503                    .map(|val| val.map(|v| PgInterval::new(0, 0, v / 1_000i64)))
504                    .collect();
505                encoder.encode_field(&value, pg_field)?;
506                Ok(())
507            }
508        },
509        DataType::Interval(interval_unit) => match interval_unit {
510            IntervalUnit::YearMonth => {
511                let value: Vec<Option<PgInterval>> = arr
512                    .as_any()
513                    .downcast_ref::<IntervalYearMonthArray>()
514                    .unwrap()
515                    .iter()
516                    .map(|val| val.map(|v| PgInterval::new(v, 0, 0)))
517                    .collect();
518                encoder.encode_field(&value, pg_field)?;
519                Ok(())
520            }
521            IntervalUnit::DayTime => {
522                let value: Vec<Option<PgInterval>> = arr
523                    .as_any()
524                    .downcast_ref::<IntervalDayTimeArray>()
525                    .unwrap()
526                    .iter()
527                    .map(|val| {
528                        val.map(|v| {
529                            let (days, millis) = IntervalDayTimeType::to_parts(v);
530                            PgInterval::new(0, days, millis as i64 * 1000i64)
531                        })
532                    })
533                    .collect();
534                encoder.encode_field(&value, pg_field)?;
535                Ok(())
536            }
537            IntervalUnit::MonthDayNano => {
538                let value: Vec<Option<PgInterval>> = arr
539                    .as_any()
540                    .downcast_ref::<IntervalMonthDayNanoArray>()
541                    .unwrap()
542                    .iter()
543                    .map(|val| {
544                        val.map(|v| {
545                            let (months, days, nanos) = IntervalMonthDayNanoType::to_parts(v);
546                            PgInterval::new(months, days, nanos / 1000i64)
547                        })
548                    })
549                    .collect();
550                encoder.encode_field(&value, pg_field)?;
551                Ok(())
552            }
553        },
554        DataType::List(_) => {
555            // Support for nested lists (list of lists)
556            // For now, convert to string representation
557            let list_array = arr.as_any().downcast_ref::<ListArray>().unwrap();
558            let value: Vec<Option<String>> = (0..list_array.len())
559                .map(|i| {
560                    if list_array.is_null(i) {
561                        None
562                    } else {
563                        // Convert nested list to string representation
564                        Some(format!("[nested_list_{i}]"))
565                    }
566                })
567                .collect();
568            encoder.encode_field(&value, pg_field)?;
569            Ok(())
570        }
571        DataType::LargeList(_) => {
572            // Support for large lists
573            let list_array = arr.as_any().downcast_ref::<LargeListArray>().unwrap();
574            let value: Vec<Option<String>> = (0..list_array.len())
575                .map(|i| {
576                    if list_array.is_null(i) {
577                        None
578                    } else {
579                        Some(format!("[large_list_{i}]"))
580                    }
581                })
582                .collect();
583            encoder.encode_field(&value, pg_field)
584        }
585        DataType::Map(_, _) => {
586            // Support for map types
587            let map_array = arr.as_any().downcast_ref::<MapArray>().unwrap();
588            let value: Vec<Option<String>> = (0..map_array.len())
589                .map(|i| {
590                    if map_array.is_null(i) {
591                        None
592                    } else {
593                        Some(format!("{{map_{i}}}"))
594                    }
595                })
596                .collect();
597            encoder.encode_field(&value, pg_field)?;
598            Ok(())
599        }
600
601        DataType::Union(_, _) => {
602            // Support for union types
603            let value: Vec<Option<String>> = (0..arr.len())
604                .map(|i| {
605                    if arr.is_null(i) {
606                        None
607                    } else {
608                        Some(format!("union_{i}"))
609                    }
610                })
611                .collect();
612            encoder.encode_field(&value, pg_field)?;
613            Ok(())
614        }
615        DataType::Dictionary(_, _) => {
616            // Support for dictionary types
617            let value: Vec<Option<String>> = (0..arr.len())
618                .map(|i| {
619                    if arr.is_null(i) {
620                        None
621                    } else {
622                        Some(format!("dict_{i}"))
623                    }
624                })
625                .collect();
626            encoder.encode_field(&value, pg_field)?;
627            Ok(())
628        }
629        // TODO: add support for more advanced types (fixed size lists, etc.)
630        list_type => Err(PgWireError::ApiError(ToSqlError::from(format!(
631            "Unsupported List Datatype {} and array {:?}",
632            list_type, arr
633        )))),
634    }
635}