|
| 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 | +} |
0 commit comments