Skip to content

Commit 0c3cfb7

Browse files
authored
fix(query): centralize conversion safety checks (#19845)
1 parent 8de2e60 commit 0c3cfb7

11 files changed

Lines changed: 1109 additions & 193 deletions

File tree

Lines changed: 359 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,359 @@
1+
// Copyright 2021 Datafuse Labs
2+
//
3+
// Licensed under the Apache License, Version 2.0 (the "License");
4+
// you may not use this file except in compliance with the License.
5+
// You may obtain a copy of the License at
6+
//
7+
// http://www.apache.org/licenses/LICENSE-2.0
8+
//
9+
// Unless required by applicable law or agreed to in writing, software
10+
// distributed under the License is distributed on an "AS IS" BASIS,
11+
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+
// See the License for the specific language governing permissions and
13+
// limitations under the License.
14+
15+
use crate::type_check::common_super_type;
16+
use crate::types::DataType;
17+
use crate::types::Decimal;
18+
use crate::types::DecimalSize;
19+
use crate::types::NumberDataType;
20+
use crate::types::i256;
21+
22+
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
23+
pub enum ConversionClass {
24+
/// Same logical type, no conversion needed.
25+
Identity,
26+
/// Conversion preserves distinct source values in the target type.
27+
LosslessInjective,
28+
/// Conversion is deterministic but can lose information or merge values.
29+
Lossy,
30+
/// Conversion semantics depend on runtime contents, for example String -> Number.
31+
ValueDependent,
32+
/// Conversion is represented by TRY_CAST or may turn failures into NULL.
33+
TryOnly,
34+
/// No supported conversion is known.
35+
Unsupported,
36+
}
37+
38+
impl ConversionClass {
39+
pub fn is_lossless_injective(&self) -> bool {
40+
matches!(self, Self::Identity | Self::LosslessInjective)
41+
}
42+
43+
pub fn is_safe_for_equality_inference(&self) -> bool {
44+
self.is_lossless_injective()
45+
}
46+
}
47+
48+
#[derive(Debug, Clone, PartialEq, Eq)]
49+
pub struct CommonTypeConversion {
50+
pub common_type: DataType,
51+
pub left: ConversionClass,
52+
pub right: ConversionClass,
53+
}
54+
55+
impl CommonTypeConversion {
56+
pub fn is_safe_for_equality_inference(&self) -> bool {
57+
is_type_safe_for_equality_inference(&self.common_type)
58+
&& self.left.is_safe_for_equality_inference()
59+
&& self.right.is_safe_for_equality_inference()
60+
}
61+
}
62+
63+
impl From<&DataType> for DataType {
64+
fn from(value: &DataType) -> Self {
65+
value.clone()
66+
}
67+
}
68+
69+
pub fn classify_conversion(src: &DataType, dest: &DataType) -> ConversionClass {
70+
if src == dest {
71+
return ConversionClass::Identity;
72+
}
73+
74+
match (src, dest) {
75+
(DataType::Null, _) => ConversionClass::LosslessInjective,
76+
(DataType::Nullable(_), DataType::Null) => ConversionClass::Lossy,
77+
(_, DataType::Null) => ConversionClass::Unsupported,
78+
79+
(DataType::Nullable(src), DataType::Nullable(dest)) => classify_conversion(src, dest),
80+
(DataType::Nullable(_), _) => ConversionClass::Unsupported,
81+
(src, DataType::Nullable(dest)) => match classify_conversion(src, dest) {
82+
ConversionClass::Identity => ConversionClass::LosslessInjective,
83+
class => class,
84+
},
85+
86+
(DataType::EmptyArray, DataType::Array(_)) => ConversionClass::LosslessInjective,
87+
(DataType::EmptyMap, DataType::Map(_)) => ConversionClass::LosslessInjective,
88+
89+
(DataType::Array(src), DataType::Array(dest))
90+
| (DataType::Map(src), DataType::Map(dest)) => classify_conversion(src, dest),
91+
(DataType::Tuple(src), DataType::Tuple(dest)) if src.len() == dest.len() => {
92+
combine_conversion_classes(
93+
src.iter()
94+
.zip(dest)
95+
.map(|(src, dest)| classify_conversion(src, dest)),
96+
)
97+
}
98+
99+
(DataType::Number(src), DataType::Number(dest)) => classify_number_conversion(*src, *dest),
100+
(DataType::Number(src), DataType::Decimal(dest)) => {
101+
if let Some(src) = src.get_decimal_properties() {
102+
match classify_decimal_conversion(src, *dest) {
103+
ConversionClass::Identity => ConversionClass::LosslessInjective,
104+
class => class,
105+
}
106+
} else {
107+
ConversionClass::Lossy
108+
}
109+
}
110+
(DataType::Decimal(src), DataType::Decimal(dest)) => {
111+
classify_decimal_conversion(*src, *dest)
112+
}
113+
(DataType::Decimal(_), DataType::Number(dest)) if dest.is_float() => ConversionClass::Lossy,
114+
(DataType::Decimal(_), DataType::Number(_)) => ConversionClass::Lossy,
115+
116+
(DataType::Boolean, DataType::String | DataType::Number(_) | DataType::Decimal(_)) => {
117+
ConversionClass::LosslessInjective
118+
}
119+
(DataType::Number(_), DataType::String) | (DataType::Decimal(_), DataType::String) => {
120+
ConversionClass::LosslessInjective
121+
}
122+
123+
(DataType::Date, DataType::Timestamp) => ConversionClass::LosslessInjective,
124+
(DataType::Timestamp, DataType::Date) => ConversionClass::Lossy,
125+
126+
(DataType::String, dest) if is_value_dependent_string_target(dest) => {
127+
ConversionClass::ValueDependent
128+
}
129+
(DataType::Variant, dest) if is_variant_try_cast_target(dest) => ConversionClass::TryOnly,
130+
131+
_ => ConversionClass::Unsupported,
132+
}
133+
}
134+
135+
pub fn common_super_type_with_conversion(
136+
left: impl Into<DataType>,
137+
right: impl Into<DataType>,
138+
) -> Option<CommonTypeConversion> {
139+
let left = left.into();
140+
let right = right.into();
141+
let common_type = common_type_for_conversion(left.clone(), right.clone())?;
142+
let left_class = classify_conversion(&left, &common_type);
143+
let right_class = classify_conversion(&right, &common_type);
144+
145+
if matches!(left_class, ConversionClass::Unsupported)
146+
|| matches!(right_class, ConversionClass::Unsupported)
147+
{
148+
return None;
149+
}
150+
151+
Some(CommonTypeConversion {
152+
common_type,
153+
left: left_class,
154+
right: right_class,
155+
})
156+
}
157+
158+
fn common_type_for_conversion(left: DataType, right: DataType) -> Option<DataType> {
159+
match (left, right) {
160+
(DataType::Null, DataType::Null) => Some(DataType::Null),
161+
(DataType::Null, ty @ DataType::Nullable(_))
162+
| (ty @ DataType::Nullable(_), DataType::Null) => Some(ty),
163+
(DataType::Null, ty) | (ty, DataType::Null) => Some(DataType::Nullable(Box::new(ty))),
164+
165+
(DataType::Nullable(left), DataType::Nullable(right)) => {
166+
Some(common_type_for_conversion(*left, *right)?.wrap_nullable())
167+
}
168+
(DataType::Nullable(left), right) => {
169+
Some(common_type_for_conversion(*left, right)?.wrap_nullable())
170+
}
171+
(left, DataType::Nullable(right)) => {
172+
Some(common_type_for_conversion(left, *right)?.wrap_nullable())
173+
}
174+
175+
(DataType::EmptyArray, ty @ DataType::Array(_))
176+
| (ty @ DataType::Array(_), DataType::EmptyArray) => Some(ty),
177+
(DataType::Array(left), DataType::Array(right)) => Some(DataType::Array(Box::new(
178+
common_type_for_conversion(*left, *right)?,
179+
))),
180+
181+
(DataType::EmptyMap, ty @ DataType::Map(_))
182+
| (ty @ DataType::Map(_), DataType::EmptyMap) => Some(ty),
183+
(DataType::Map(left), DataType::Map(right)) => Some(DataType::Map(Box::new(
184+
common_type_for_conversion(*left, *right)?,
185+
))),
186+
187+
(DataType::Tuple(left), DataType::Tuple(right)) if left.len() == right.len() => {
188+
let fields = left
189+
.into_iter()
190+
.zip(right)
191+
.map(|(left, right)| common_type_for_conversion(left, right))
192+
.collect::<Option<Vec<_>>>()?;
193+
Some(DataType::Tuple(fields))
194+
}
195+
196+
(DataType::Number(left), DataType::Number(right)) => Some(number_common_type(left, right)),
197+
198+
(left, right) => {
199+
if let Some(common_type) = common_super_type(left.clone(), right.clone(), &[]) {
200+
return Some(common_type);
201+
}
202+
203+
match (left, right) {
204+
(DataType::Date, DataType::Timestamp) | (DataType::Timestamp, DataType::Date) => {
205+
Some(DataType::Timestamp)
206+
}
207+
208+
(DataType::String, ty) if is_value_dependent_string_target(&ty) => Some(ty),
209+
(ty, DataType::String) if is_value_dependent_string_target(&ty) => Some(ty),
210+
211+
(DataType::Variant, ty) if is_variant_try_cast_target(&ty) => {
212+
Some(ty.wrap_nullable())
213+
}
214+
(ty, DataType::Variant) if is_variant_try_cast_target(&ty) => {
215+
Some(ty.wrap_nullable())
216+
}
217+
218+
_ => None,
219+
}
220+
}
221+
}
222+
}
223+
224+
fn classify_number_conversion(src: NumberDataType, dest: NumberDataType) -> ConversionClass {
225+
if src == dest {
226+
return ConversionClass::Identity;
227+
}
228+
229+
match (src.is_integer(), dest.is_integer()) {
230+
(true, true) if src.can_lossless_cast_to(dest) => ConversionClass::LosslessInjective,
231+
(true, true) => ConversionClass::Lossy,
232+
(false, false) if src.can_lossless_cast_to(dest) => ConversionClass::LosslessInjective,
233+
(false, false) => ConversionClass::Lossy,
234+
(true, false) | (false, true) => ConversionClass::Lossy,
235+
}
236+
}
237+
238+
fn classify_decimal_conversion(src: DecimalSize, dest: DecimalSize) -> ConversionClass {
239+
if src == dest {
240+
return ConversionClass::Identity;
241+
}
242+
243+
if src.scale() <= dest.scale() && src.leading_digits() <= dest.leading_digits() {
244+
ConversionClass::LosslessInjective
245+
} else {
246+
ConversionClass::Lossy
247+
}
248+
}
249+
250+
fn number_common_type(left: NumberDataType, right: NumberDataType) -> DataType {
251+
if left == right {
252+
return DataType::Number(left);
253+
}
254+
255+
if left.is_integer() && right.is_integer() {
256+
if left.can_lossless_cast_to(right) {
257+
return DataType::Number(right);
258+
}
259+
if right.can_lossless_cast_to(left) {
260+
return DataType::Number(left);
261+
}
262+
263+
let left = left.get_decimal_properties().unwrap();
264+
let right = right.get_decimal_properties().unwrap();
265+
return DataType::Decimal(decimal_common_size(left, right));
266+
}
267+
268+
if left.is_float() && right.is_float() {
269+
if left.can_lossless_cast_to(right) {
270+
return DataType::Number(right);
271+
}
272+
return DataType::Number(left);
273+
}
274+
275+
if left.is_float() {
276+
DataType::Number(left)
277+
} else {
278+
DataType::Number(right)
279+
}
280+
}
281+
282+
fn decimal_common_size(left: DecimalSize, right: DecimalSize) -> DecimalSize {
283+
let scale = left.scale().max(right.scale());
284+
let precision = scale + left.leading_digits().max(right.leading_digits());
285+
let precision =
286+
if left.precision() <= i128::MAX_PRECISION && right.precision() <= i128::MAX_PRECISION {
287+
precision.min(i128::MAX_PRECISION)
288+
} else {
289+
precision.min(i256::MAX_PRECISION)
290+
};
291+
292+
DecimalSize::new_unchecked(precision, scale)
293+
}
294+
295+
fn combine_conversion_classes(
296+
classes: impl IntoIterator<Item = ConversionClass>,
297+
) -> ConversionClass {
298+
let mut result = ConversionClass::Identity;
299+
for class in classes {
300+
result = match (result, class) {
301+
(ConversionClass::Unsupported, _) | (_, ConversionClass::Unsupported) => {
302+
ConversionClass::Unsupported
303+
}
304+
(ConversionClass::TryOnly, _) | (_, ConversionClass::TryOnly) => {
305+
ConversionClass::TryOnly
306+
}
307+
(ConversionClass::ValueDependent, _) | (_, ConversionClass::ValueDependent) => {
308+
ConversionClass::ValueDependent
309+
}
310+
(ConversionClass::Lossy, _) | (_, ConversionClass::Lossy) => ConversionClass::Lossy,
311+
(ConversionClass::LosslessInjective, _) | (_, ConversionClass::LosslessInjective) => {
312+
ConversionClass::LosslessInjective
313+
}
314+
(ConversionClass::Identity, ConversionClass::Identity) => ConversionClass::Identity,
315+
};
316+
}
317+
result
318+
}
319+
320+
fn is_value_dependent_string_target(ty: &DataType) -> bool {
321+
match ty {
322+
DataType::Nullable(ty) => is_value_dependent_string_target(ty),
323+
DataType::Number(_)
324+
| DataType::Decimal(_)
325+
| DataType::Boolean
326+
| DataType::Date
327+
| DataType::Timestamp
328+
| DataType::TimestampTz
329+
| DataType::Interval => true,
330+
_ => false,
331+
}
332+
}
333+
334+
fn is_variant_try_cast_target(ty: &DataType) -> bool {
335+
match ty {
336+
DataType::Nullable(ty) => is_variant_try_cast_target(ty),
337+
DataType::Boolean
338+
| DataType::Date
339+
| DataType::Timestamp
340+
| DataType::String
341+
| DataType::Number(_) => true,
342+
_ => false,
343+
}
344+
}
345+
346+
fn is_type_safe_for_equality_inference(ty: &DataType) -> bool {
347+
!matches!(
348+
ty.remove_nullable(),
349+
DataType::Map(_)
350+
| DataType::EmptyMap
351+
| DataType::Binary
352+
| DataType::Geometry
353+
| DataType::Geography
354+
| DataType::Vector(_)
355+
| DataType::Opaque(_)
356+
| DataType::Generic(_)
357+
| DataType::StageLocation
358+
)
359+
}

src/query/expression/src/lib.rs

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -56,6 +56,7 @@ mod block;
5656
pub mod aggregate;
5757
mod block_vec;
5858
mod constant_folder;
59+
pub mod conversion;
5960
pub mod converts;
6061
mod evaluator;
6162
mod expression;

0 commit comments

Comments
 (0)