Skip to main content

lance_datafusion/
expr.rs

1// SPDX-License-Identifier: Apache-2.0
2// SPDX-FileCopyrightText: Copyright The Lance Authors
3
4//! Utilities for working with datafusion expressions
5
6use std::sync::Arc;
7
8use arrow::compute::cast;
9use arrow_array::{ArrayRef, cast::AsArray};
10use arrow_schema::{DataType, TimeUnit};
11use datafusion_common::ScalarValue;
12
13const MS_PER_DAY: i64 = 86400000;
14
15// This is slightly tedious but when we convert expressions from SQL strings to logical
16// datafusion expressions there is no type coercion that happens.  In other words "x = 7"
17// will always yield "x = 7_u64" regardless of the type of the column "x".  As a result, we
18// need to do that literal coercion ourselves.
19pub fn safe_coerce_scalar(value: &ScalarValue, ty: &DataType) -> Option<ScalarValue> {
20    // A dictionary target coerces the value to the dictionary's value type and
21    // re-wraps it as a dictionary literal. Only an untyped `ScalarValue::Null`
22    // keeps its untyped form, matching the behavior for all other targets; a
23    // *typed* null (e.g. `Utf8(None)`) is coerced and wrapped like any other
24    // value so it produces a `Dictionary(..)` literal that matches the column.
25    if let DataType::Dictionary(key_type, value_type) = ty {
26        if matches!(value, ScalarValue::Null) {
27            return Some(value.clone());
28        }
29        let inner = safe_coerce_scalar(value, value_type)?;
30        return Some(ScalarValue::Dictionary(key_type.clone(), Box::new(inner)));
31    }
32    match value {
33        ScalarValue::Int8(val) => match ty {
34            DataType::Int8 => Some(value.clone()),
35            DataType::Int16 => val.map(|v| ScalarValue::Int16(Some(i16::from(v)))),
36            DataType::Int32 => val.map(|v| ScalarValue::Int32(Some(i32::from(v)))),
37            DataType::Int64 => val.map(|v| ScalarValue::Int64(Some(i64::from(v)))),
38            DataType::UInt8 => {
39                val.and_then(|v| u8::try_from(v).map(|v| ScalarValue::UInt8(Some(v))).ok())
40            }
41            DataType::UInt16 => {
42                val.and_then(|v| u16::try_from(v).map(|v| ScalarValue::UInt16(Some(v))).ok())
43            }
44            DataType::UInt32 => {
45                val.and_then(|v| u32::try_from(v).map(|v| ScalarValue::UInt32(Some(v))).ok())
46            }
47            DataType::UInt64 => {
48                val.and_then(|v| u64::try_from(v).map(|v| ScalarValue::UInt64(Some(v))).ok())
49            }
50            DataType::Float32 => val.map(|v| ScalarValue::Float32(Some(f32::from(v)))),
51            DataType::Float64 => val.map(|v| ScalarValue::Float64(Some(f64::from(v)))),
52            _ => None,
53        },
54        ScalarValue::Int16(val) => match ty {
55            DataType::Int8 => {
56                val.and_then(|v| i8::try_from(v).map(|v| ScalarValue::Int8(Some(v))).ok())
57            }
58            DataType::Int16 => Some(value.clone()),
59            DataType::Int32 => val.map(|v| ScalarValue::Int32(Some(i32::from(v)))),
60            DataType::Int64 => val.map(|v| ScalarValue::Int64(Some(i64::from(v)))),
61            DataType::UInt8 => {
62                val.and_then(|v| u8::try_from(v).map(|v| ScalarValue::UInt8(Some(v))).ok())
63            }
64            DataType::UInt16 => {
65                val.and_then(|v| u16::try_from(v).map(|v| ScalarValue::UInt16(Some(v))).ok())
66            }
67            DataType::UInt32 => {
68                val.and_then(|v| u32::try_from(v).map(|v| ScalarValue::UInt32(Some(v))).ok())
69            }
70            DataType::UInt64 => {
71                val.and_then(|v| u64::try_from(v).map(|v| ScalarValue::UInt64(Some(v))).ok())
72            }
73            DataType::Float32 => val.map(|v| ScalarValue::Float32(Some(f32::from(v)))),
74            DataType::Float64 => val.map(|v| ScalarValue::Float64(Some(f64::from(v)))),
75            _ => None,
76        },
77        ScalarValue::Int32(val) => match ty {
78            DataType::Int8 => {
79                val.and_then(|v| i8::try_from(v).map(|v| ScalarValue::Int8(Some(v))).ok())
80            }
81            DataType::Int16 => {
82                val.and_then(|v| i16::try_from(v).map(|v| ScalarValue::Int16(Some(v))).ok())
83            }
84            DataType::Int32 => Some(value.clone()),
85            DataType::Int64 => val.map(|v| ScalarValue::Int64(Some(i64::from(v)))),
86            DataType::UInt8 => {
87                val.and_then(|v| u8::try_from(v).map(|v| ScalarValue::UInt8(Some(v))).ok())
88            }
89            DataType::UInt16 => {
90                val.and_then(|v| u16::try_from(v).map(|v| ScalarValue::UInt16(Some(v))).ok())
91            }
92            DataType::UInt32 => {
93                val.and_then(|v| u32::try_from(v).map(|v| ScalarValue::UInt32(Some(v))).ok())
94            }
95            DataType::UInt64 => {
96                val.and_then(|v| u64::try_from(v).map(|v| ScalarValue::UInt64(Some(v))).ok())
97            }
98            // These conversions are inherently lossy as the full range of i32 cannot
99            // be represented in f32.  However, there is no f32::TryFrom(i32) and its not
100            // clear users would want that anyways
101            DataType::Float32 => val.map(|v| ScalarValue::Float32(Some(v as f32))),
102            DataType::Float64 => val.map(|v| ScalarValue::Float64(Some(v as f64))),
103            _ => None,
104        },
105        ScalarValue::Int64(val) => match ty {
106            DataType::Int8 => {
107                val.and_then(|v| i8::try_from(v).map(|v| ScalarValue::Int8(Some(v))).ok())
108            }
109            DataType::Int16 => {
110                val.and_then(|v| i16::try_from(v).map(|v| ScalarValue::Int16(Some(v))).ok())
111            }
112            DataType::Int32 => {
113                val.and_then(|v| i32::try_from(v).map(|v| ScalarValue::Int32(Some(v))).ok())
114            }
115            DataType::Int64 => Some(value.clone()),
116            DataType::UInt8 => {
117                val.and_then(|v| u8::try_from(v).map(|v| ScalarValue::UInt8(Some(v))).ok())
118            }
119            DataType::UInt16 => {
120                val.and_then(|v| u16::try_from(v).map(|v| ScalarValue::UInt16(Some(v))).ok())
121            }
122            DataType::UInt32 => {
123                val.and_then(|v| u32::try_from(v).map(|v| ScalarValue::UInt32(Some(v))).ok())
124            }
125            DataType::UInt64 => {
126                val.and_then(|v| u64::try_from(v).map(|v| ScalarValue::UInt64(Some(v))).ok())
127            }
128            // See above warning about lossy float conversion
129            DataType::Float32 => val.map(|v| ScalarValue::Float32(Some(v as f32))),
130            DataType::Float64 => val.map(|v| ScalarValue::Float64(Some(v as f64))),
131            DataType::Decimal128(_, _) | DataType::Decimal256(_, _) => value.cast_to(ty).ok(),
132            DataType::Time32(TimeUnit::Second) => val.and_then(|v| {
133                i32::try_from(v)
134                    .ok()
135                    .map(|v| ScalarValue::Time32Second(Some(v)))
136            }),
137            DataType::Time32(TimeUnit::Millisecond) => val.and_then(|v| {
138                i32::try_from(v)
139                    .ok()
140                    .map(|v| ScalarValue::Time32Millisecond(Some(v)))
141            }),
142            _ => None,
143        },
144        ScalarValue::UInt8(val) => match ty {
145            DataType::Int8 => {
146                val.and_then(|v| i8::try_from(v).map(|v| ScalarValue::Int8(Some(v))).ok())
147            }
148            DataType::Int16 => val.map(|v| ScalarValue::Int16(Some(v.into()))),
149            DataType::Int32 => val.map(|v| ScalarValue::Int32(Some(v.into()))),
150            DataType::Int64 => val.map(|v| ScalarValue::Int64(Some(v.into()))),
151            DataType::UInt8 => Some(value.clone()),
152            DataType::UInt16 => val.map(|v| ScalarValue::UInt16(Some(u16::from(v)))),
153            DataType::UInt32 => val.map(|v| ScalarValue::UInt32(Some(u32::from(v)))),
154            DataType::UInt64 => val.map(|v| ScalarValue::UInt64(Some(u64::from(v)))),
155            DataType::Float32 => val.map(|v| ScalarValue::Float32(Some(f32::from(v)))),
156            DataType::Float64 => val.map(|v| ScalarValue::Float64(Some(f64::from(v)))),
157            _ => None,
158        },
159        ScalarValue::UInt16(val) => match ty {
160            DataType::Int8 => {
161                val.and_then(|v| i8::try_from(v).map(|v| ScalarValue::Int8(Some(v))).ok())
162            }
163            DataType::Int16 => {
164                val.and_then(|v| i16::try_from(v).map(|v| ScalarValue::Int16(Some(v))).ok())
165            }
166            DataType::Int32 => val.map(|v| ScalarValue::Int32(Some(v.into()))),
167            DataType::Int64 => val.map(|v| ScalarValue::Int64(Some(v.into()))),
168            DataType::UInt8 => {
169                val.and_then(|v| u8::try_from(v).map(|v| ScalarValue::UInt8(Some(v))).ok())
170            }
171            DataType::UInt16 => Some(value.clone()),
172            DataType::UInt32 => val.map(|v| ScalarValue::UInt32(Some(u32::from(v)))),
173            DataType::UInt64 => val.map(|v| ScalarValue::UInt64(Some(u64::from(v)))),
174            DataType::Float32 => val.map(|v| ScalarValue::Float32(Some(f32::from(v)))),
175            DataType::Float64 => val.map(|v| ScalarValue::Float64(Some(f64::from(v)))),
176            _ => None,
177        },
178        ScalarValue::UInt32(val) => match ty {
179            DataType::Int8 => {
180                val.and_then(|v| i8::try_from(v).map(|v| ScalarValue::Int8(Some(v))).ok())
181            }
182            DataType::Int16 => {
183                val.and_then(|v| i16::try_from(v).map(|v| ScalarValue::Int16(Some(v))).ok())
184            }
185            DataType::Int32 => {
186                val.and_then(|v| i32::try_from(v).map(|v| ScalarValue::Int32(Some(v))).ok())
187            }
188            DataType::Int64 => val.map(|v| ScalarValue::Int64(Some(v.into()))),
189            DataType::UInt8 => {
190                val.and_then(|v| u8::try_from(v).map(|v| ScalarValue::UInt8(Some(v))).ok())
191            }
192            DataType::UInt16 => {
193                val.and_then(|v| u16::try_from(v).map(|v| ScalarValue::UInt16(Some(v))).ok())
194            }
195            DataType::UInt32 => Some(value.clone()),
196            DataType::UInt64 => val.map(|v| ScalarValue::UInt64(Some(u64::from(v)))),
197            // See above warning about lossy float conversion
198            DataType::Float32 => val.map(|v| ScalarValue::Float32(Some(v as f32))),
199            DataType::Float64 => val.map(|v| ScalarValue::Float64(Some(v as f64))),
200            _ => None,
201        },
202        ScalarValue::UInt64(val) => match ty {
203            DataType::Int8 => {
204                val.and_then(|v| i8::try_from(v).map(|v| ScalarValue::Int8(Some(v))).ok())
205            }
206            DataType::Int16 => {
207                val.and_then(|v| i16::try_from(v).map(|v| ScalarValue::Int16(Some(v))).ok())
208            }
209            DataType::Int32 => {
210                val.and_then(|v| i32::try_from(v).map(|v| ScalarValue::Int32(Some(v))).ok())
211            }
212            DataType::Int64 => {
213                val.and_then(|v| i64::try_from(v).map(|v| ScalarValue::Int64(Some(v))).ok())
214            }
215            DataType::UInt8 => {
216                val.and_then(|v| u8::try_from(v).map(|v| ScalarValue::UInt8(Some(v))).ok())
217            }
218            DataType::UInt16 => {
219                val.and_then(|v| u16::try_from(v).map(|v| ScalarValue::UInt16(Some(v))).ok())
220            }
221            DataType::UInt32 => {
222                val.and_then(|v| u32::try_from(v).map(|v| ScalarValue::UInt32(Some(v))).ok())
223            }
224            DataType::UInt64 => Some(value.clone()),
225            // See above warning about lossy float conversion
226            DataType::Float32 => val.map(|v| ScalarValue::Float32(Some(v as f32))),
227            DataType::Float64 => val.map(|v| ScalarValue::Float64(Some(v as f64))),
228            _ => None,
229        },
230        ScalarValue::Float32(val) => match ty {
231            DataType::Float32 => Some(value.clone()),
232            DataType::Float64 => val.map(|v| ScalarValue::Float64(Some(f64::from(v)))),
233            _ => None,
234        },
235        ScalarValue::Float64(val) => match ty {
236            DataType::Float32 => val.map(|v| ScalarValue::Float32(Some(v as f32))),
237            DataType::Float64 => Some(value.clone()),
238            _ => None,
239        },
240        ScalarValue::Utf8(val) => match ty {
241            DataType::Utf8 => Some(value.clone()),
242            DataType::LargeUtf8 => Some(ScalarValue::LargeUtf8(val.clone())),
243            DataType::Utf8View => Some(ScalarValue::Utf8View(val.clone())),
244            _ => None,
245        },
246        ScalarValue::LargeUtf8(val) => match ty {
247            DataType::Utf8 => Some(ScalarValue::Utf8(val.clone())),
248            DataType::LargeUtf8 => Some(value.clone()),
249            DataType::Utf8View => Some(ScalarValue::Utf8View(val.clone())),
250            _ => None,
251        },
252        ScalarValue::Utf8View(val) => match ty {
253            DataType::Utf8 => Some(ScalarValue::Utf8(val.clone())),
254            DataType::LargeUtf8 => Some(ScalarValue::LargeUtf8(val.clone())),
255            DataType::Utf8View => Some(value.clone()),
256            _ => None,
257        },
258        ScalarValue::Boolean(_) => match ty {
259            DataType::Boolean => Some(value.clone()),
260            _ => None,
261        },
262        ScalarValue::Null => Some(value.clone()),
263        ScalarValue::List(values) => {
264            let values = values.clone() as ArrayRef;
265            let new_values = cast(&values, ty).ok()?;
266            match ty {
267                DataType::List(_) => {
268                    Some(ScalarValue::List(Arc::new(new_values.as_list().clone())))
269                }
270                DataType::LargeList(_) => Some(ScalarValue::LargeList(Arc::new(
271                    new_values.as_list().clone(),
272                ))),
273                DataType::FixedSizeList(_, _) => Some(ScalarValue::FixedSizeList(Arc::new(
274                    new_values.as_fixed_size_list().clone(),
275                ))),
276                _ => None,
277            }
278        }
279        ScalarValue::TimestampSecond(seconds, _) => match ty {
280            DataType::Timestamp(TimeUnit::Second, _) => Some(value.clone()),
281            DataType::Timestamp(TimeUnit::Millisecond, tz) => seconds
282                .and_then(|v| v.checked_mul(1000))
283                .map(|val| ScalarValue::TimestampMillisecond(Some(val), tz.clone())),
284            DataType::Timestamp(TimeUnit::Microsecond, tz) => seconds
285                .and_then(|v| v.checked_mul(1000000))
286                .map(|val| ScalarValue::TimestampMicrosecond(Some(val), tz.clone())),
287            DataType::Timestamp(TimeUnit::Nanosecond, tz) => seconds
288                .and_then(|v| v.checked_mul(1000000000))
289                .map(|val| ScalarValue::TimestampNanosecond(Some(val), tz.clone())),
290            _ => None,
291        },
292        ScalarValue::TimestampMillisecond(millis, _) => match ty {
293            DataType::Timestamp(TimeUnit::Second, tz) => {
294                millis.map(|val| ScalarValue::TimestampSecond(Some(val / 1000), tz.clone()))
295            }
296            DataType::Timestamp(TimeUnit::Millisecond, _) => Some(value.clone()),
297            DataType::Timestamp(TimeUnit::Microsecond, tz) => millis
298                .and_then(|v| v.checked_mul(1000))
299                .map(|val| ScalarValue::TimestampMicrosecond(Some(val), tz.clone())),
300            DataType::Timestamp(TimeUnit::Nanosecond, tz) => millis
301                .and_then(|v| v.checked_mul(1000000))
302                .map(|val| ScalarValue::TimestampNanosecond(Some(val), tz.clone())),
303            _ => None,
304        },
305        ScalarValue::TimestampMicrosecond(micros, _) => match ty {
306            DataType::Timestamp(TimeUnit::Second, tz) => {
307                micros.map(|val| ScalarValue::TimestampSecond(Some(val / 1000000), tz.clone()))
308            }
309            DataType::Timestamp(TimeUnit::Millisecond, tz) => {
310                micros.map(|val| ScalarValue::TimestampMillisecond(Some(val / 1000), tz.clone()))
311            }
312            DataType::Timestamp(TimeUnit::Microsecond, _) => Some(value.clone()),
313            DataType::Timestamp(TimeUnit::Nanosecond, tz) => micros
314                .and_then(|v| v.checked_mul(1000))
315                .map(|val| ScalarValue::TimestampNanosecond(Some(val), tz.clone())),
316            _ => None,
317        },
318        ScalarValue::TimestampNanosecond(nanos, _) => {
319            match ty {
320                DataType::Timestamp(TimeUnit::Second, tz) => nanos
321                    .map(|val| ScalarValue::TimestampSecond(Some(val / 1000000000), tz.clone())),
322                DataType::Timestamp(TimeUnit::Millisecond, tz) => nanos
323                    .map(|val| ScalarValue::TimestampMillisecond(Some(val / 1000000), tz.clone())),
324                DataType::Timestamp(TimeUnit::Microsecond, tz) => {
325                    nanos.map(|val| ScalarValue::TimestampMicrosecond(Some(val / 1000), tz.clone()))
326                }
327                DataType::Timestamp(TimeUnit::Nanosecond, _) => Some(value.clone()),
328                _ => None,
329            }
330        }
331        ScalarValue::Date32(ticks) => match ty {
332            DataType::Date32 => Some(value.clone()),
333            DataType::Date64 => Some(ScalarValue::Date64(
334                ticks.map(|v| i64::from(v) * MS_PER_DAY),
335            )),
336            _ => None,
337        },
338        ScalarValue::Date64(ticks) => match ty {
339            DataType::Date32 => Some(ScalarValue::Date32(ticks.map(|v| (v / MS_PER_DAY) as i32))),
340            DataType::Date64 => Some(value.clone()),
341            _ => None,
342        },
343        ScalarValue::Time32Second(seconds) => {
344            match ty {
345                DataType::Time32(TimeUnit::Second) => Some(value.clone()),
346                DataType::Time32(TimeUnit::Millisecond) => {
347                    seconds.map(|val| ScalarValue::Time32Millisecond(Some(val * 1000)))
348                }
349                DataType::Time64(TimeUnit::Microsecond) => seconds
350                    .map(|val| ScalarValue::Time64Microsecond(Some(i64::from(val) * 1000000))),
351                DataType::Time64(TimeUnit::Nanosecond) => seconds
352                    .map(|val| ScalarValue::Time64Nanosecond(Some(i64::from(val) * 1000000000))),
353                _ => None,
354            }
355        }
356        ScalarValue::Time32Millisecond(millis) => match ty {
357            DataType::Time32(TimeUnit::Second) => {
358                millis.map(|val| ScalarValue::Time32Second(Some(val / 1000)))
359            }
360            DataType::Time32(TimeUnit::Millisecond) => Some(value.clone()),
361            DataType::Time64(TimeUnit::Microsecond) => {
362                millis.map(|val| ScalarValue::Time64Microsecond(Some(i64::from(val) * 1000)))
363            }
364            DataType::Time64(TimeUnit::Nanosecond) => {
365                millis.map(|val| ScalarValue::Time64Nanosecond(Some(i64::from(val) * 1000000)))
366            }
367            _ => None,
368        },
369        ScalarValue::Time64Microsecond(micros) => match ty {
370            DataType::Time32(TimeUnit::Second) => {
371                micros.map(|val| ScalarValue::Time32Second(Some((val / 1000000) as i32)))
372            }
373            DataType::Time32(TimeUnit::Millisecond) => {
374                micros.map(|val| ScalarValue::Time32Millisecond(Some((val / 1000) as i32)))
375            }
376            DataType::Time64(TimeUnit::Microsecond) => Some(value.clone()),
377            DataType::Time64(TimeUnit::Nanosecond) => {
378                micros.map(|val| ScalarValue::Time64Nanosecond(Some(val * 1000)))
379            }
380            _ => None,
381        },
382        ScalarValue::Time64Nanosecond(nanos) => match ty {
383            DataType::Time32(TimeUnit::Second) => {
384                nanos.map(|val| ScalarValue::Time32Second(Some((val / 1000000000) as i32)))
385            }
386            DataType::Time32(TimeUnit::Millisecond) => {
387                nanos.map(|val| ScalarValue::Time32Millisecond(Some((val / 1000000) as i32)))
388            }
389            DataType::Time64(TimeUnit::Microsecond) => {
390                nanos.map(|val| ScalarValue::Time64Microsecond(Some(val / 1000)))
391            }
392            DataType::Time64(TimeUnit::Nanosecond) => Some(value.clone()),
393            _ => None,
394        },
395        ScalarValue::LargeList(values) => {
396            let values = values.clone() as ArrayRef;
397            let new_values = cast(&values, ty).ok()?;
398            match ty {
399                DataType::List(_) => {
400                    Some(ScalarValue::List(Arc::new(new_values.as_list().clone())))
401                }
402                DataType::LargeList(_) => Some(ScalarValue::LargeList(Arc::new(
403                    new_values.as_list().clone(),
404                ))),
405                DataType::FixedSizeList(_, _) => Some(ScalarValue::FixedSizeList(Arc::new(
406                    new_values.as_fixed_size_list().clone(),
407                ))),
408                _ => None,
409            }
410        }
411        ScalarValue::FixedSizeList(values) => {
412            let values = values.clone() as ArrayRef;
413            let new_values = cast(&values, ty).ok()?;
414            match ty {
415                DataType::List(_) => {
416                    Some(ScalarValue::List(Arc::new(new_values.as_list().clone())))
417                }
418                DataType::LargeList(_) => Some(ScalarValue::LargeList(Arc::new(
419                    new_values.as_list().clone(),
420                ))),
421                DataType::FixedSizeList(_, _) => Some(ScalarValue::FixedSizeList(Arc::new(
422                    new_values.as_fixed_size_list().clone(),
423                ))),
424                _ => None,
425            }
426        }
427        ScalarValue::FixedSizeBinary(len, value) => match ty {
428            DataType::FixedSizeBinary(len2) => {
429                if len == len2 {
430                    Some(ScalarValue::FixedSizeBinary(*len, value.clone()))
431                } else {
432                    None
433                }
434            }
435            DataType::Binary => Some(ScalarValue::Binary(value.clone())),
436            _ => None,
437        },
438        ScalarValue::Binary(value) => match ty {
439            DataType::Binary => Some(ScalarValue::Binary(value.clone())),
440            DataType::LargeBinary => Some(ScalarValue::LargeBinary(value.clone())),
441            DataType::BinaryView => Some(ScalarValue::BinaryView(value.clone())),
442            DataType::FixedSizeBinary(len) => {
443                if let Some(value) = value {
444                    if value.len() == *len as usize {
445                        Some(ScalarValue::FixedSizeBinary(*len, Some(value.clone())))
446                    } else {
447                        None
448                    }
449                } else {
450                    None
451                }
452            }
453            _ => None,
454        },
455        ScalarValue::BinaryView(val) => match ty {
456            DataType::Binary => Some(ScalarValue::Binary(val.clone())),
457            DataType::LargeBinary => Some(ScalarValue::LargeBinary(val.clone())),
458            DataType::BinaryView => Some(value.clone()),
459            _ => None,
460        },
461        ScalarValue::LargeBinary(_) => match ty {
462            DataType::LargeBinary => Some(value.clone()),
463            _ => None,
464        },
465        ScalarValue::Decimal128(_, _, _) => match ty {
466            DataType::Decimal128(_, _) => value.cast_to(ty).ok(),
467            _ => None,
468        },
469        ScalarValue::Decimal256(_, _, _) => match ty {
470            DataType::Decimal256(_, _) => value.cast_to(ty).ok(),
471            _ => None,
472        },
473        ScalarValue::DurationSecond(_)
474        | ScalarValue::DurationMillisecond(_)
475        | ScalarValue::DurationMicrosecond(_)
476        | ScalarValue::DurationNanosecond(_) => match ty {
477            DataType::Duration(_) => value.cast_to(ty).ok(),
478            _ => None,
479        },
480        // A dictionary-encoded literal (e.g. produced by DataFusion's dictionary
481        // cast in the scalar-index path) coerces by unwrapping its underlying value.
482        ScalarValue::Dictionary(_, inner) => safe_coerce_scalar(inner, ty),
483        _ => None,
484    }
485}
486
487#[cfg(test)]
488mod tests {
489    use arrow::datatypes::i256;
490
491    use super::*;
492
493    #[test]
494    fn test_temporal_coerce() {
495        assert_eq!(
496            safe_coerce_scalar(
497                &ScalarValue::Int64(Some(5)),
498                &DataType::Time32(TimeUnit::Second),
499            ),
500            Some(ScalarValue::Time32Second(Some(5)))
501        );
502        assert_eq!(
503            safe_coerce_scalar(
504                &ScalarValue::Int64(Some(5000)),
505                &DataType::Time32(TimeUnit::Millisecond),
506            ),
507            Some(ScalarValue::Time32Millisecond(Some(5000)))
508        );
509        assert_eq!(
510            safe_coerce_scalar(
511                &ScalarValue::Int64(Some(i64::MAX)),
512                &DataType::Time32(TimeUnit::Second),
513            ),
514            None
515        );
516
517        // Conversion from timestamps in one resolution to timestamps in another resolution is allowed
518        // s->s
519        assert_eq!(
520            safe_coerce_scalar(
521                &ScalarValue::TimestampSecond(Some(5), None),
522                &DataType::Timestamp(TimeUnit::Second, None),
523            ),
524            Some(ScalarValue::TimestampSecond(Some(5), None))
525        );
526        // s->ms
527        assert_eq!(
528            safe_coerce_scalar(
529                &ScalarValue::TimestampSecond(Some(5), None),
530                &DataType::Timestamp(TimeUnit::Millisecond, None),
531            ),
532            Some(ScalarValue::TimestampMillisecond(Some(5000), None))
533        );
534        // s->us
535        assert_eq!(
536            safe_coerce_scalar(
537                &ScalarValue::TimestampSecond(Some(5), None),
538                &DataType::Timestamp(TimeUnit::Microsecond, None),
539            ),
540            Some(ScalarValue::TimestampMicrosecond(Some(5000000), None))
541        );
542        // s->ns
543        assert_eq!(
544            safe_coerce_scalar(
545                &ScalarValue::TimestampSecond(Some(5), None),
546                &DataType::Timestamp(TimeUnit::Nanosecond, None),
547            ),
548            Some(ScalarValue::TimestampNanosecond(Some(5000000000), None))
549        );
550        // ms->s
551        assert_eq!(
552            safe_coerce_scalar(
553                &ScalarValue::TimestampMillisecond(Some(5000), None),
554                &DataType::Timestamp(TimeUnit::Second, None),
555            ),
556            Some(ScalarValue::TimestampSecond(Some(5), None))
557        );
558        // ms->ms
559        assert_eq!(
560            safe_coerce_scalar(
561                &ScalarValue::TimestampMillisecond(Some(5000), None),
562                &DataType::Timestamp(TimeUnit::Millisecond, None),
563            ),
564            Some(ScalarValue::TimestampMillisecond(Some(5000), None))
565        );
566        // ms->us
567        assert_eq!(
568            safe_coerce_scalar(
569                &ScalarValue::TimestampMillisecond(Some(5000), None),
570                &DataType::Timestamp(TimeUnit::Microsecond, None),
571            ),
572            Some(ScalarValue::TimestampMicrosecond(Some(5000000), None))
573        );
574        // ms->ns
575        assert_eq!(
576            safe_coerce_scalar(
577                &ScalarValue::TimestampMillisecond(Some(5000), None),
578                &DataType::Timestamp(TimeUnit::Nanosecond, None),
579            ),
580            Some(ScalarValue::TimestampNanosecond(Some(5000000000), None))
581        );
582        // us->s
583        assert_eq!(
584            safe_coerce_scalar(
585                &ScalarValue::TimestampMicrosecond(Some(5000000), None),
586                &DataType::Timestamp(TimeUnit::Second, None),
587            ),
588            Some(ScalarValue::TimestampSecond(Some(5), None))
589        );
590        // us->ms
591        assert_eq!(
592            safe_coerce_scalar(
593                &ScalarValue::TimestampMicrosecond(Some(5000000), None),
594                &DataType::Timestamp(TimeUnit::Millisecond, None),
595            ),
596            Some(ScalarValue::TimestampMillisecond(Some(5000), None))
597        );
598        // us->us
599        assert_eq!(
600            safe_coerce_scalar(
601                &ScalarValue::TimestampMicrosecond(Some(5000000), None),
602                &DataType::Timestamp(TimeUnit::Microsecond, None),
603            ),
604            Some(ScalarValue::TimestampMicrosecond(Some(5000000), None))
605        );
606        // us->ns
607        assert_eq!(
608            safe_coerce_scalar(
609                &ScalarValue::TimestampMicrosecond(Some(5000000), None),
610                &DataType::Timestamp(TimeUnit::Nanosecond, None),
611            ),
612            Some(ScalarValue::TimestampNanosecond(Some(5000000000), None))
613        );
614        // ns->s
615        assert_eq!(
616            safe_coerce_scalar(
617                &ScalarValue::TimestampNanosecond(Some(5000000000), None),
618                &DataType::Timestamp(TimeUnit::Second, None),
619            ),
620            Some(ScalarValue::TimestampSecond(Some(5), None))
621        );
622        // ns->ms
623        assert_eq!(
624            safe_coerce_scalar(
625                &ScalarValue::TimestampNanosecond(Some(5000000000), None),
626                &DataType::Timestamp(TimeUnit::Millisecond, None),
627            ),
628            Some(ScalarValue::TimestampMillisecond(Some(5000), None))
629        );
630        // ns->us
631        assert_eq!(
632            safe_coerce_scalar(
633                &ScalarValue::TimestampNanosecond(Some(5000000000), None),
634                &DataType::Timestamp(TimeUnit::Microsecond, None),
635            ),
636            Some(ScalarValue::TimestampMicrosecond(Some(5000000), None))
637        );
638        // ns->ns
639        assert_eq!(
640            safe_coerce_scalar(
641                &ScalarValue::TimestampNanosecond(Some(5000000000), None),
642                &DataType::Timestamp(TimeUnit::Nanosecond, None),
643            ),
644            Some(ScalarValue::TimestampNanosecond(Some(5000000000), None))
645        );
646        // Precision loss on coercion is allowed (truncation)
647        // ns->s
648        assert_eq!(
649            safe_coerce_scalar(
650                &ScalarValue::TimestampNanosecond(Some(5987654321), None),
651                &DataType::Timestamp(TimeUnit::Second, None),
652            ),
653            Some(ScalarValue::TimestampSecond(Some(5), None))
654        );
655        // Conversions from date-32 to date-64 is allowed
656        assert_eq!(
657            safe_coerce_scalar(&ScalarValue::Date32(Some(5)), &DataType::Date32,),
658            Some(ScalarValue::Date32(Some(5)))
659        );
660        assert_eq!(
661            safe_coerce_scalar(&ScalarValue::Date32(Some(5)), &DataType::Date64,),
662            Some(ScalarValue::Date64(Some(5 * MS_PER_DAY)))
663        );
664        assert_eq!(
665            safe_coerce_scalar(
666                &ScalarValue::Date64(Some(5 * MS_PER_DAY)),
667                &DataType::Date32,
668            ),
669            Some(ScalarValue::Date32(Some(5)))
670        );
671        assert_eq!(
672            safe_coerce_scalar(&ScalarValue::Date64(Some(5)), &DataType::Date64,),
673            Some(ScalarValue::Date64(Some(5)))
674        );
675        // Time-32 to time-64 (and within time-32 and time-64) is allowed
676        assert_eq!(
677            safe_coerce_scalar(
678                &ScalarValue::Time32Second(Some(5)),
679                &DataType::Time32(TimeUnit::Second),
680            ),
681            Some(ScalarValue::Time32Second(Some(5)))
682        );
683        assert_eq!(
684            safe_coerce_scalar(
685                &ScalarValue::Time32Second(Some(5)),
686                &DataType::Time32(TimeUnit::Millisecond),
687            ),
688            Some(ScalarValue::Time32Millisecond(Some(5000)))
689        );
690        assert_eq!(
691            safe_coerce_scalar(
692                &ScalarValue::Time32Second(Some(5)),
693                &DataType::Time64(TimeUnit::Microsecond),
694            ),
695            Some(ScalarValue::Time64Microsecond(Some(5000000)))
696        );
697        assert_eq!(
698            safe_coerce_scalar(
699                &ScalarValue::Time32Second(Some(5)),
700                &DataType::Time64(TimeUnit::Nanosecond),
701            ),
702            Some(ScalarValue::Time64Nanosecond(Some(5000000000)))
703        );
704        assert_eq!(
705            safe_coerce_scalar(
706                &ScalarValue::Time32Millisecond(Some(5000)),
707                &DataType::Time32(TimeUnit::Second),
708            ),
709            Some(ScalarValue::Time32Second(Some(5)))
710        );
711        assert_eq!(
712            safe_coerce_scalar(
713                &ScalarValue::Time32Millisecond(Some(5000)),
714                &DataType::Time32(TimeUnit::Millisecond),
715            ),
716            Some(ScalarValue::Time32Millisecond(Some(5000)))
717        );
718        assert_eq!(
719            safe_coerce_scalar(
720                &ScalarValue::Time32Millisecond(Some(5000)),
721                &DataType::Time64(TimeUnit::Microsecond),
722            ),
723            Some(ScalarValue::Time64Microsecond(Some(5000000)))
724        );
725        assert_eq!(
726            safe_coerce_scalar(
727                &ScalarValue::Time32Millisecond(Some(5000)),
728                &DataType::Time64(TimeUnit::Nanosecond),
729            ),
730            Some(ScalarValue::Time64Nanosecond(Some(5000000000)))
731        );
732        assert_eq!(
733            safe_coerce_scalar(
734                &ScalarValue::Time64Microsecond(Some(5000000)),
735                &DataType::Time32(TimeUnit::Second),
736            ),
737            Some(ScalarValue::Time32Second(Some(5)))
738        );
739        assert_eq!(
740            safe_coerce_scalar(
741                &ScalarValue::Time64Microsecond(Some(5000000)),
742                &DataType::Time32(TimeUnit::Millisecond),
743            ),
744            Some(ScalarValue::Time32Millisecond(Some(5000)))
745        );
746        assert_eq!(
747            safe_coerce_scalar(
748                &ScalarValue::Time64Microsecond(Some(5000000)),
749                &DataType::Time64(TimeUnit::Microsecond),
750            ),
751            Some(ScalarValue::Time64Microsecond(Some(5000000)))
752        );
753        assert_eq!(
754            safe_coerce_scalar(
755                &ScalarValue::Time64Microsecond(Some(5000000)),
756                &DataType::Time64(TimeUnit::Nanosecond),
757            ),
758            Some(ScalarValue::Time64Nanosecond(Some(5000000000)))
759        );
760        assert_eq!(
761            safe_coerce_scalar(
762                &ScalarValue::Time64Nanosecond(Some(5000000000)),
763                &DataType::Time32(TimeUnit::Second),
764            ),
765            Some(ScalarValue::Time32Second(Some(5)))
766        );
767        assert_eq!(
768            safe_coerce_scalar(
769                &ScalarValue::Time64Nanosecond(Some(5000000000)),
770                &DataType::Time32(TimeUnit::Millisecond),
771            ),
772            Some(ScalarValue::Time32Millisecond(Some(5000)))
773        );
774        assert_eq!(
775            safe_coerce_scalar(
776                &ScalarValue::Time64Nanosecond(Some(5000000000)),
777                &DataType::Time64(TimeUnit::Microsecond),
778            ),
779            Some(ScalarValue::Time64Microsecond(Some(5000000)))
780        );
781        assert_eq!(
782            safe_coerce_scalar(
783                &ScalarValue::Time64Nanosecond(Some(5000000000)),
784                &DataType::Time64(TimeUnit::Nanosecond),
785            ),
786            Some(ScalarValue::Time64Nanosecond(Some(5000000000)))
787        );
788        assert_eq!(
789            safe_coerce_scalar(
790                &ScalarValue::DurationNanosecond(Some(2_000_000)),
791                &DataType::Duration(TimeUnit::Millisecond),
792            ),
793            Some(ScalarValue::DurationMillisecond(Some(2)))
794        );
795    }
796
797    #[test]
798    fn test_string_view_coerce() {
799        // Utf8 <-> Utf8View
800        assert_eq!(
801            safe_coerce_scalar(&ScalarValue::Utf8(Some("hi".into())), &DataType::Utf8View),
802            Some(ScalarValue::Utf8View(Some("hi".into())))
803        );
804        assert_eq!(
805            safe_coerce_scalar(&ScalarValue::Utf8View(Some("hi".into())), &DataType::Utf8),
806            Some(ScalarValue::Utf8(Some("hi".into())))
807        );
808        assert_eq!(
809            safe_coerce_scalar(
810                &ScalarValue::Utf8View(Some("hi".into())),
811                &DataType::LargeUtf8
812            ),
813            Some(ScalarValue::LargeUtf8(Some("hi".into())))
814        );
815        assert_eq!(
816            safe_coerce_scalar(
817                &ScalarValue::LargeUtf8(Some("hi".into())),
818                &DataType::Utf8View
819            ),
820            Some(ScalarValue::Utf8View(Some("hi".into())))
821        );
822        // identity
823        assert_eq!(
824            safe_coerce_scalar(
825                &ScalarValue::Utf8View(Some("hi".into())),
826                &DataType::Utf8View
827            ),
828            Some(ScalarValue::Utf8View(Some("hi".into())))
829        );
830        // Binary <-> BinaryView
831        assert_eq!(
832            safe_coerce_scalar(
833                &ScalarValue::Binary(Some(vec![1, 2, 3])),
834                &DataType::BinaryView
835            ),
836            Some(ScalarValue::BinaryView(Some(vec![1, 2, 3])))
837        );
838        assert_eq!(
839            safe_coerce_scalar(
840                &ScalarValue::BinaryView(Some(vec![1, 2, 3])),
841                &DataType::Binary
842            ),
843            Some(ScalarValue::Binary(Some(vec![1, 2, 3])))
844        );
845        assert_eq!(
846            safe_coerce_scalar(
847                &ScalarValue::BinaryView(Some(vec![1, 2, 3])),
848                &DataType::BinaryView
849            ),
850            Some(ScalarValue::BinaryView(Some(vec![1, 2, 3])))
851        );
852        assert_eq!(
853            safe_coerce_scalar(
854                &ScalarValue::LargeBinary(Some(vec![1, 2, 3])),
855                &DataType::LargeBinary
856            ),
857            Some(ScalarValue::LargeBinary(Some(vec![1, 2, 3])))
858        );
859    }
860
861    #[test]
862    fn test_decimal_coerce() {
863        assert_eq!(
864            safe_coerce_scalar(
865                &ScalarValue::Decimal128(Some(2), 10, 0),
866                &DataType::Decimal128(12, 2),
867            ),
868            Some(ScalarValue::Decimal128(Some(200), 12, 2))
869        );
870        assert_eq!(
871            safe_coerce_scalar(
872                &ScalarValue::Decimal256(Some(i256::from_i128(2)), 76, 0),
873                &DataType::Decimal256(76, 2),
874            ),
875            Some(ScalarValue::Decimal256(Some(i256::from_i128(200)), 76, 2))
876        );
877    }
878
879    #[test]
880    fn test_dictionary_coerce() {
881        let dict_ty = DataType::Dictionary(Box::new(DataType::Int16), Box::new(DataType::Utf8));
882
883        // A string literal coerces to a dictionary target by wrapping the
884        // coerced value in a dictionary scalar.
885        assert_eq!(
886            safe_coerce_scalar(&ScalarValue::Utf8(Some("com".to_string())), &dict_ty),
887            Some(ScalarValue::Dictionary(
888                Box::new(DataType::Int16),
889                Box::new(ScalarValue::Utf8(Some("com".to_string()))),
890            ))
891        );
892
893        // The inner value is coerced through to the dictionary value type, so a
894        // LargeUtf8 literal lands as a Utf8 value inside the dictionary.
895        assert_eq!(
896            safe_coerce_scalar(&ScalarValue::LargeUtf8(Some("com".to_string())), &dict_ty),
897            Some(ScalarValue::Dictionary(
898                Box::new(DataType::Int16),
899                Box::new(ScalarValue::Utf8(Some("com".to_string()))),
900            ))
901        );
902
903        // A dictionary literal round-trips back to its value type.
904        assert_eq!(
905            safe_coerce_scalar(
906                &ScalarValue::Dictionary(
907                    Box::new(DataType::Int16),
908                    Box::new(ScalarValue::Utf8(Some("com".to_string()))),
909                ),
910                &DataType::Utf8,
911            ),
912            Some(ScalarValue::Utf8(Some("com".to_string())))
913        );
914
915        // A dictionary literal coerces to a dictionary target, adopting the
916        // target's key type.
917        assert_eq!(
918            safe_coerce_scalar(
919                &ScalarValue::Dictionary(
920                    Box::new(DataType::Int32),
921                    Box::new(ScalarValue::Utf8(Some("com".to_string()))),
922                ),
923                &dict_ty,
924            ),
925            Some(ScalarValue::Dictionary(
926                Box::new(DataType::Int16),
927                Box::new(ScalarValue::Utf8(Some("com".to_string()))),
928            ))
929        );
930
931        // An untyped null keeps its untyped form for a dictionary target, just
932        // like for every other target type.
933        assert_eq!(
934            safe_coerce_scalar(&ScalarValue::Null, &dict_ty),
935            Some(ScalarValue::Null)
936        );
937
938        // A *typed* null (e.g. an API-built `Utf8(None)` literal, or an IN value
939        // already typed as Utf8) is still wrapped in the dictionary type so it
940        // matches the dictionary column. Returning a bare `Utf8(None)` here would
941        // leave `resolve_value` with a literal whose type does not line up with
942        // the column, breaking planning/evaluation the same way non-null strings
943        // used to break.
944        assert_eq!(
945            safe_coerce_scalar(&ScalarValue::Utf8(None), &dict_ty),
946            Some(ScalarValue::Dictionary(
947                Box::new(DataType::Int16),
948                Box::new(ScalarValue::Utf8(None)),
949            ))
950        );
951
952        // The inner null is coerced through to the dictionary value type as well,
953        // so a LargeUtf8 typed null lands as a Utf8 null inside the dictionary.
954        assert_eq!(
955            safe_coerce_scalar(&ScalarValue::LargeUtf8(None), &dict_ty),
956            Some(ScalarValue::Dictionary(
957                Box::new(DataType::Int16),
958                Box::new(ScalarValue::Utf8(None)),
959            ))
960        );
961
962        // A value that cannot be coerced to the dictionary value type fails.
963        assert_eq!(
964            safe_coerce_scalar(
965                &ScalarValue::Utf8(Some("com".to_string())),
966                &DataType::Dictionary(Box::new(DataType::Int16), Box::new(DataType::Int32)),
967            ),
968            None
969        );
970    }
971}