diff --git a/minarrow-py/src/convert.rs b/minarrow-py/src/convert.rs index 8f0b233..6040f7f 100644 --- a/minarrow-py/src/convert.rs +++ b/minarrow-py/src/convert.rs @@ -628,11 +628,11 @@ pub fn scalar_to_py(py: Python<'_>, scalar: Scalar) -> PyResult> { #[cfg(feature = "large_string")] Scalar::String64(v) => v.into_py_any(py), #[cfg(feature = "decimal")] - Scalar::Decimal32(v, scale) => decimal_scalar_to_py(py, v as i128, scale), + Scalar::Decimal32(v, _, scale) => decimal_scalar_to_py(py, v as i128, scale), #[cfg(feature = "decimal")] - Scalar::Decimal64(v, scale) => decimal_scalar_to_py(py, v as i128, scale), + Scalar::Decimal64(v, _, scale) => decimal_scalar_to_py(py, v as i128, scale), #[cfg(feature = "decimal")] - Scalar::Decimal128(v, scale) => decimal_scalar_to_py(py, v, scale), + Scalar::Decimal128(v, _, scale) => decimal_scalar_to_py(py, v, scale), #[cfg(feature = "datetime")] Scalar::Datetime32(v) => v.into_py_any(py), #[cfg(feature = "datetime")] @@ -668,11 +668,11 @@ pub fn scalar_repr(scalar: &Scalar) -> String { #[cfg(feature = "large_string")] Scalar::String64(v) => format!("\"{}\"", v), #[cfg(feature = "decimal")] - Scalar::Decimal32(v, s) => format_decimal_string(*v as i128, *s), + Scalar::Decimal32(v, _, s) => format_decimal_string(*v as i128, *s), #[cfg(feature = "decimal")] - Scalar::Decimal64(v, s) => format_decimal_string(*v as i128, *s), + Scalar::Decimal64(v, _, s) => format_decimal_string(*v as i128, *s), #[cfg(feature = "decimal")] - Scalar::Decimal128(v, s) => format_decimal_string(*v, *s), + Scalar::Decimal128(v, _, s) => format_decimal_string(*v, *s), #[cfg(feature = "datetime")] Scalar::Datetime32(v) => v.to_string(), #[cfg(feature = "datetime")] diff --git a/src/conversions.rs b/src/conversions.rs index dbe4606..e713d55 100644 --- a/src/conversions.rs +++ b/src/conversions.rs @@ -1438,16 +1438,16 @@ impl From for Array { #[cfg(feature = "datetime")] Interval => Array::from_int32(IntegerArray::from_slice(&[0i32])), #[cfg(feature = "decimal")] - Decimal32(v, s) => Array::from_decimal32( - crate::DecimalArray::from_slice(&[v], 0, s), + Decimal32(v, p, s) => Array::from_decimal32( + crate::DecimalArray::from_slice(&[v], p, s), ), #[cfg(feature = "decimal")] - Decimal64(v, s) => Array::from_decimal64( - crate::DecimalArray::from_slice(&[v], 0, s), + Decimal64(v, p, s) => Array::from_decimal64( + crate::DecimalArray::from_slice(&[v], p, s), ), #[cfg(feature = "decimal")] - Decimal128(v, s) => Array::from_decimal128( - crate::DecimalArray::from_slice(&[v], 0, s), + Decimal128(v, p, s) => Array::from_decimal128( + crate::DecimalArray::from_slice(&[v], p, s), ), } } diff --git a/src/enums/array.rs b/src/enums/array.rs index a7d387a..e5bf87d 100644 --- a/src/enums/array.rs +++ b/src/enums/array.rs @@ -2535,11 +2535,11 @@ impl Array { NumericArray::Float32(a) => Some(Scalar::Float32(a.data[idx])), NumericArray::Float64(a) => Some(Scalar::Float64(a.data[idx])), #[cfg(feature = "decimal")] - NumericArray::Decimal32(a) => Some(Scalar::Decimal32(a.data[idx], a.scale)), + NumericArray::Decimal32(a) => Some(Scalar::Decimal32(a.data[idx], a.precision, a.scale)), #[cfg(feature = "decimal")] - NumericArray::Decimal64(a) => Some(Scalar::Decimal64(a.data[idx], a.scale)), + NumericArray::Decimal64(a) => Some(Scalar::Decimal64(a.data[idx], a.precision, a.scale)), #[cfg(feature = "decimal")] - NumericArray::Decimal128(a) => Some(Scalar::Decimal128(a.data[idx], a.scale)), + NumericArray::Decimal128(a) => Some(Scalar::Decimal128(a.data[idx], a.precision, a.scale)), NumericArray::Null => Some(Scalar::Null), }, Array::TextArray(text) => match text { @@ -3139,13 +3139,14 @@ impl Array { )) } #[cfg(feature = "decimal")] - Scalar::Decimal32(_, scale) => { + Scalar::Decimal32(_, precision, scale) => { + let precision = *precision; let scale = *scale; let mut data = Vec64::::with_capacity(scalars.len()); let mut mask = Bitmask::new_set_all(scalars.len(), true); for (i, s) in scalars.iter().enumerate() { match s { - Scalar::Decimal32(v, _) => data.push(*v), + Scalar::Decimal32(v, _, _) => data.push(*v), Scalar::Null => { data.push(0); mask.set(i, false); @@ -3157,18 +3158,19 @@ impl Array { Array::from_decimal32(crate::DecimalArray::new( data, if has_nulls { Some(mask) } else { None }, - 0, + precision, scale, )) } #[cfg(feature = "decimal")] - Scalar::Decimal64(_, scale) => { + Scalar::Decimal64(_, precision, scale) => { + let precision = *precision; let scale = *scale; let mut data = Vec64::::with_capacity(scalars.len()); let mut mask = Bitmask::new_set_all(scalars.len(), true); for (i, s) in scalars.iter().enumerate() { match s { - Scalar::Decimal64(v, _) => data.push(*v), + Scalar::Decimal64(v, _, _) => data.push(*v), Scalar::Null => { data.push(0); mask.set(i, false); @@ -3180,18 +3182,19 @@ impl Array { Array::from_decimal64(crate::DecimalArray::new( data, if has_nulls { Some(mask) } else { None }, - 0, + precision, scale, )) } #[cfg(feature = "decimal")] - Scalar::Decimal128(_, scale) => { + Scalar::Decimal128(_, precision, scale) => { + let precision = *precision; let scale = *scale; let mut data = Vec64::::with_capacity(scalars.len()); let mut mask = Bitmask::new_set_all(scalars.len(), true); for (i, s) in scalars.iter().enumerate() { match s { - Scalar::Decimal128(v, _) => data.push(*v), + Scalar::Decimal128(v, _, _) => data.push(*v), Scalar::Null => { data.push(0); mask.set(i, false); @@ -3203,7 +3206,7 @@ impl Array { Array::from_decimal128(crate::DecimalArray::new( data, if has_nulls { Some(mask) } else { None }, - 0, + precision, scale, )) } @@ -7393,8 +7396,9 @@ mod arr_macro_extensions_tests { let arr = Array::from_decimal32(d); let s = arr.get_scalar(0).unwrap(); match s { - crate::Scalar::Decimal32(v, scale) => { + crate::Scalar::Decimal32(v, precision, scale) => { assert_eq!(v, 12345); + assert_eq!(precision, 9); assert_eq!(scale, 2); } _ => panic!("expected Scalar::Decimal32"), @@ -7404,21 +7408,21 @@ mod arr_macro_extensions_tests { #[cfg(feature = "scalar_type")] #[test] fn scalar_decimal_display() { - let s = crate::Scalar::Decimal64(12345, 2); + let s = crate::Scalar::Decimal64(12345, 18, 2); assert_eq!(format!("{}", s), "123.45"); } #[cfg(feature = "scalar_type")] #[test] fn scalar_decimal_zero_scale() { - let s = crate::Scalar::Decimal128(42, 0); + let s = crate::Scalar::Decimal128(42, 38, 0); assert_eq!(format!("{}", s), "42"); } #[cfg(feature = "scalar_type")] #[test] fn scalar_decimal_negative_scale() { - let s = crate::Scalar::Decimal32(5, -2); + let s = crate::Scalar::Decimal32(5, 9, -2); assert_eq!(format!("{}", s), "500"); } diff --git a/src/enums/scalar.rs b/src/enums/scalar.rs index 6ed21cb..e497d82 100644 --- a/src/enums/scalar.rs +++ b/src/enums/scalar.rs @@ -103,13 +103,15 @@ pub enum Scalar { // Floats Float32(f32), Float64(f64), - // Decimals - unscaled integer value + scale + // Decimals - unscaled integer value, precision, scale. The scalar carries + // both components of its `ArrowType` so it round-trips to its column type + // without loss. #[cfg(feature = "decimal")] - Decimal32(i32, i8), + Decimal32(i32, u8, i8), #[cfg(feature = "decimal")] - Decimal64(i64, i8), + Decimal64(i64, u8, i8), #[cfg(feature = "decimal")] - Decimal128(i128, i8), + Decimal128(i128, u8, i8), // String strings String32(String), #[cfg(feature = "large_string")] @@ -143,11 +145,11 @@ impl Display for Scalar { Scalar::Float32(v) => Display::fmt(v, f), Scalar::Float64(v) => Display::fmt(v, f), #[cfg(feature = "decimal")] - Scalar::Decimal32(v, scale) => write!(f, "{}", format_decimal_scalar(*v as i128, *scale)), + Scalar::Decimal32(v, _, scale) => write!(f, "{}", format_decimal_scalar(*v as i128, *scale)), #[cfg(feature = "decimal")] - Scalar::Decimal64(v, scale) => write!(f, "{}", format_decimal_scalar(*v as i128, *scale)), + Scalar::Decimal64(v, _, scale) => write!(f, "{}", format_decimal_scalar(*v as i128, *scale)), #[cfg(feature = "decimal")] - Scalar::Decimal128(v, scale) => write!(f, "{}", format_decimal_scalar(*v, *scale)), + Scalar::Decimal128(v, _, scale) => write!(f, "{}", format_decimal_scalar(*v, *scale)), Scalar::String32(v) => f.write_str(v), #[cfg(feature = "large_string")] Scalar::String64(v) => f.write_str(v), @@ -186,11 +188,11 @@ impl Scalar { Scalar::Float32(v) => *v != 0.0, Scalar::Float64(v) => *v != 0.0, #[cfg(feature = "decimal")] - Scalar::Decimal32(v, _) => *v != 0, + Scalar::Decimal32(v, _, _) => *v != 0, #[cfg(feature = "decimal")] - Scalar::Decimal64(v, _) => *v != 0, + Scalar::Decimal64(v, _, _) => *v != 0, #[cfg(feature = "decimal")] - Scalar::Decimal128(v, _) => *v != 0, + Scalar::Decimal128(v, _, _) => *v != 0, Scalar::Null => panic!("Cannot convert Null to bool"), Scalar::String32(s) => { let s = s.trim(); @@ -257,11 +259,11 @@ impl Scalar { Scalar::Float32(v) => i8::try_from(*v as i32).expect("f32 out of range for i8"), Scalar::Float64(v) => i8::try_from(*v as i32).expect("f64 out of range for i8"), #[cfg(feature = "decimal")] - Scalar::Decimal32(v, _) => i8::try_from(*v).expect("Decimal32 out of range for i8"), + Scalar::Decimal32(v, _, _) => i8::try_from(*v).expect("Decimal32 out of range for i8"), #[cfg(feature = "decimal")] - Scalar::Decimal64(v, _) => i8::try_from(*v).expect("Decimal64 out of range for i8"), + Scalar::Decimal64(v, _, _) => i8::try_from(*v).expect("Decimal64 out of range for i8"), #[cfg(feature = "decimal")] - Scalar::Decimal128(v, _) => i8::try_from(*v).expect("Decimal128 out of range for i8"), + Scalar::Decimal128(v, _, _) => i8::try_from(*v).expect("Decimal128 out of range for i8"), Scalar::Null => panic!("Cannot convert Null to i8"), Scalar::String32(s) => s.parse::().expect("Cannot parse string as i8"), #[cfg(feature = "large_string")] @@ -301,11 +303,11 @@ impl Scalar { Scalar::Float32(v) => i16::try_from(*v as i32).expect("f32 out of range for i16"), Scalar::Float64(v) => i16::try_from(*v as i32).expect("f64 out of range for i16"), #[cfg(feature = "decimal")] - Scalar::Decimal32(v, _) => i16::try_from(*v).expect("Decimal32 out of range for i16"), + Scalar::Decimal32(v, _, _) => i16::try_from(*v).expect("Decimal32 out of range for i16"), #[cfg(feature = "decimal")] - Scalar::Decimal64(v, _) => i16::try_from(*v).expect("Decimal64 out of range for i16"), + Scalar::Decimal64(v, _, _) => i16::try_from(*v).expect("Decimal64 out of range for i16"), #[cfg(feature = "decimal")] - Scalar::Decimal128(v, _) => i16::try_from(*v).expect("Decimal128 out of range for i16"), + Scalar::Decimal128(v, _, _) => i16::try_from(*v).expect("Decimal128 out of range for i16"), Scalar::Null => panic!("Cannot convert Null to i16"), Scalar::String32(s) => s.parse::().expect("Cannot parse string as i16"), #[cfg(feature = "large_string")] @@ -348,11 +350,11 @@ impl Scalar { Scalar::Float32(v) => *v as i32, Scalar::Float64(v) => *v as i32, #[cfg(feature = "decimal")] - Scalar::Decimal32(v, _) => *v, + Scalar::Decimal32(v, _, _) => *v, #[cfg(feature = "decimal")] - Scalar::Decimal64(v, _) => i32::try_from(*v).expect("Decimal64 out of range for i32"), + Scalar::Decimal64(v, _, _) => i32::try_from(*v).expect("Decimal64 out of range for i32"), #[cfg(feature = "decimal")] - Scalar::Decimal128(v, _) => i32::try_from(*v).expect("Decimal128 out of range for i32"), + Scalar::Decimal128(v, _, _) => i32::try_from(*v).expect("Decimal128 out of range for i32"), Scalar::Null => panic!("Cannot convert Null to i32"), Scalar::String32(s) => s.parse::().expect("Cannot parse string as i32"), #[cfg(feature = "large_string")] @@ -401,11 +403,11 @@ impl Scalar { Scalar::Float32(v) => *v as i64, Scalar::Float64(v) => *v as i64, #[cfg(feature = "decimal")] - Scalar::Decimal32(v, _) => *v as i64, + Scalar::Decimal32(v, _, _) => *v as i64, #[cfg(feature = "decimal")] - Scalar::Decimal64(v, _) => *v, + Scalar::Decimal64(v, _, _) => *v, #[cfg(feature = "decimal")] - Scalar::Decimal128(v, _) => i64::try_from(*v).expect("Decimal128 out of range for i64"), + Scalar::Decimal128(v, _, _) => i64::try_from(*v).expect("Decimal128 out of range for i64"), Scalar::Null => panic!("Cannot convert Null to i64"), Scalar::String32(s) => s.parse::().expect("Cannot parse string as i64"), #[cfg(feature = "large_string")] @@ -448,11 +450,11 @@ impl Scalar { Scalar::Float32(v) => u8::try_from(*v as i32).expect("f32 out of range for u8"), Scalar::Float64(v) => u8::try_from(*v as i32).expect("f64 out of range for u8"), #[cfg(feature = "decimal")] - Scalar::Decimal32(v, _) => u8::try_from(*v).expect("Decimal32 out of range for u8"), + Scalar::Decimal32(v, _, _) => u8::try_from(*v).expect("Decimal32 out of range for u8"), #[cfg(feature = "decimal")] - Scalar::Decimal64(v, _) => u8::try_from(*v).expect("Decimal64 out of range for u8"), + Scalar::Decimal64(v, _, _) => u8::try_from(*v).expect("Decimal64 out of range for u8"), #[cfg(feature = "decimal")] - Scalar::Decimal128(v, _) => u8::try_from(*v).expect("Decimal128 out of range for u8"), + Scalar::Decimal128(v, _, _) => u8::try_from(*v).expect("Decimal128 out of range for u8"), Scalar::Null => panic!("Cannot convert Null to u8"), Scalar::String32(s) => s.parse::().expect("Cannot parse string as u8"), #[cfg(feature = "large_string")] @@ -495,11 +497,11 @@ impl Scalar { Scalar::Float32(v) => u16::try_from(*v as i32).expect("f32 out of range for u16"), Scalar::Float64(v) => u16::try_from(*v as i32).expect("f64 out of range for u16"), #[cfg(feature = "decimal")] - Scalar::Decimal32(v, _) => u16::try_from(*v).expect("Decimal32 out of range for u16"), + Scalar::Decimal32(v, _, _) => u16::try_from(*v).expect("Decimal32 out of range for u16"), #[cfg(feature = "decimal")] - Scalar::Decimal64(v, _) => u16::try_from(*v).expect("Decimal64 out of range for u16"), + Scalar::Decimal64(v, _, _) => u16::try_from(*v).expect("Decimal64 out of range for u16"), #[cfg(feature = "decimal")] - Scalar::Decimal128(v, _) => u16::try_from(*v).expect("Decimal128 out of range for u16"), + Scalar::Decimal128(v, _, _) => u16::try_from(*v).expect("Decimal128 out of range for u16"), Scalar::Null => panic!("Cannot convert Null to u16"), Scalar::String32(s) => s.parse::().expect("Cannot parse string as u16"), #[cfg(feature = "large_string")] @@ -542,11 +544,11 @@ impl Scalar { Scalar::Float32(v) => *v as u32, Scalar::Float64(v) => *v as u32, #[cfg(feature = "decimal")] - Scalar::Decimal32(v, _) => u32::try_from(*v).expect("Decimal32 out of range for u32"), + Scalar::Decimal32(v, _, _) => u32::try_from(*v).expect("Decimal32 out of range for u32"), #[cfg(feature = "decimal")] - Scalar::Decimal64(v, _) => u32::try_from(*v).expect("Decimal64 out of range for u32"), + Scalar::Decimal64(v, _, _) => u32::try_from(*v).expect("Decimal64 out of range for u32"), #[cfg(feature = "decimal")] - Scalar::Decimal128(v, _) => u32::try_from(*v).expect("Decimal128 out of range for u32"), + Scalar::Decimal128(v, _, _) => u32::try_from(*v).expect("Decimal128 out of range for u32"), Scalar::Null => panic!("Cannot convert Null to u32"), Scalar::String32(s) => s.parse::().expect("Cannot parse string as u32"), #[cfg(feature = "large_string")] @@ -625,15 +627,15 @@ impl Scalar { } } #[cfg(feature = "decimal")] - Scalar::Decimal32(v, _) => { + Scalar::Decimal32(v, _, _) => { if *v >= 0 { *v as u64 } else { panic!("Decimal32 out of range for u64") } }, #[cfg(feature = "decimal")] - Scalar::Decimal64(v, _) => { + Scalar::Decimal64(v, _, _) => { if *v >= 0 { *v as u64 } else { panic!("Decimal64 out of range for u64") } }, #[cfg(feature = "decimal")] - Scalar::Decimal128(v, _) => u64::try_from(*v).expect("Decimal128 out of range for u64"), + Scalar::Decimal128(v, _, _) => u64::try_from(*v).expect("Decimal128 out of range for u64"), Scalar::Null => panic!("Cannot convert Null to u64"), Scalar::String32(s) => s.parse::().expect("Cannot parse string as u64"), #[cfg(feature = "large_string")] @@ -676,11 +678,11 @@ impl Scalar { Scalar::Float32(v) => *v, Scalar::Float64(v) => *v as f32, #[cfg(feature = "decimal")] - Scalar::Decimal32(v, s) => *v as f32 / 10f32.powi(*s as i32), + Scalar::Decimal32(v, _, s) => *v as f32 / 10f32.powi(*s as i32), #[cfg(feature = "decimal")] - Scalar::Decimal64(v, s) => *v as f32 / 10f32.powi(*s as i32), + Scalar::Decimal64(v, _, s) => *v as f32 / 10f32.powi(*s as i32), #[cfg(feature = "decimal")] - Scalar::Decimal128(v, s) => *v as f32 / 10f32.powi(*s as i32), + Scalar::Decimal128(v, _, s) => *v as f32 / 10f32.powi(*s as i32), Scalar::Boolean(v) => { if *v { 1.0 @@ -723,11 +725,11 @@ impl Scalar { Scalar::Float32(v) => *v as f64, Scalar::Float64(v) => *v, #[cfg(feature = "decimal")] - Scalar::Decimal32(v, s) => *v as f64 / 10f64.powi(*s as i32), + Scalar::Decimal32(v, _, s) => *v as f64 / 10f64.powi(*s as i32), #[cfg(feature = "decimal")] - Scalar::Decimal64(v, s) => *v as f64 / 10f64.powi(*s as i32), + Scalar::Decimal64(v, _, s) => *v as f64 / 10f64.powi(*s as i32), #[cfg(feature = "decimal")] - Scalar::Decimal128(v, s) => *v as f64 / 10f64.powi(*s as i32), + Scalar::Decimal128(v, _, s) => *v as f64 / 10f64.powi(*s as i32), Scalar::Boolean(v) => { if *v { 1.0 @@ -774,11 +776,11 @@ impl Scalar { Scalar::Float32(v) => v.to_string(), Scalar::Float64(v) => v.to_string(), #[cfg(feature = "decimal")] - Scalar::Decimal32(v, s) => format_decimal_scalar(*v as i128, *s), + Scalar::Decimal32(v, _, s) => format_decimal_scalar(*v as i128, *s), #[cfg(feature = "decimal")] - Scalar::Decimal64(v, s) => format_decimal_scalar(*v as i128, *s), + Scalar::Decimal64(v, _, s) => format_decimal_scalar(*v as i128, *s), #[cfg(feature = "decimal")] - Scalar::Decimal128(v, s) => format_decimal_scalar(*v, *s), + Scalar::Decimal128(v, _, s) => format_decimal_scalar(*v, *s), Scalar::Null => panic!("Cannot convert Null to String"), #[cfg(feature = "datetime")] Scalar::Datetime32(v) => v.to_string(), @@ -849,7 +851,7 @@ impl Scalar { } } #[cfg(feature = "decimal")] - Scalar::Decimal32(_, _) | Scalar::Decimal64(_, _) | Scalar::Decimal128(_, _) => panic!("Cannot convert Decimal to dt32"), + Scalar::Decimal32(_, _, _) | Scalar::Decimal64(_, _, _) | Scalar::Decimal128(_, _, _) => panic!("Cannot convert Decimal to dt32"), Scalar::String32(s) => s.parse::().expect("Cannot parse string as dt32"), #[cfg(feature = "large_string")] Scalar::String64(s) => s.parse::().expect("Cannot parse string as dt32"), @@ -938,7 +940,7 @@ impl Scalar { } } #[cfg(feature = "decimal")] - Scalar::Decimal32(_, _) | Scalar::Decimal64(_, _) | Scalar::Decimal128(_, _) => panic!("Cannot convert Decimal to dt64"), + Scalar::Decimal32(_, _, _) | Scalar::Decimal64(_, _, _) | Scalar::Decimal128(_, _, _) => panic!("Cannot convert Decimal to dt64"), Scalar::String32(s) => s.parse::().expect("Cannot parse string as dt64"), #[cfg(feature = "large_string")] Scalar::String64(s) => s.parse::().expect("Cannot parse string as dt64"), @@ -979,11 +981,11 @@ impl Scalar { Scalar::Float32(v) => Some(*v != 0.0), Scalar::Float64(v) => Some(*v != 0.0), #[cfg(feature = "decimal")] - Scalar::Decimal32(v, _) => Some(*v != 0), + Scalar::Decimal32(v, _, _) => Some(*v != 0), #[cfg(feature = "decimal")] - Scalar::Decimal64(v, _) => Some(*v != 0), + Scalar::Decimal64(v, _, _) => Some(*v != 0), #[cfg(feature = "decimal")] - Scalar::Decimal128(v, _) => Some(*v != 0), + Scalar::Decimal128(v, _, _) => Some(*v != 0), Scalar::Null => None, Scalar::String32(s) => { let s = s.trim(); @@ -1047,11 +1049,11 @@ impl Scalar { Scalar::Float32(v) => i8::try_from(*v as i32).ok(), Scalar::Float64(v) => i8::try_from(*v as i32).ok(), #[cfg(feature = "decimal")] - Scalar::Decimal32(v, _) => i8::try_from(*v).ok(), + Scalar::Decimal32(v, _, _) => i8::try_from(*v).ok(), #[cfg(feature = "decimal")] - Scalar::Decimal64(v, _) => i8::try_from(*v).ok(), + Scalar::Decimal64(v, _, _) => i8::try_from(*v).ok(), #[cfg(feature = "decimal")] - Scalar::Decimal128(v, _) => i8::try_from(*v).ok(), + Scalar::Decimal128(v, _, _) => i8::try_from(*v).ok(), Scalar::Null => None, Scalar::String32(s) => s.parse::().ok(), #[cfg(feature = "large_string")] @@ -1082,11 +1084,11 @@ impl Scalar { Scalar::Float32(v) => i16::try_from(*v as i32).ok(), Scalar::Float64(v) => i16::try_from(*v as i32).ok(), #[cfg(feature = "decimal")] - Scalar::Decimal32(v, _) => i16::try_from(*v).ok(), + Scalar::Decimal32(v, _, _) => i16::try_from(*v).ok(), #[cfg(feature = "decimal")] - Scalar::Decimal64(v, _) => i16::try_from(*v).ok(), + Scalar::Decimal64(v, _, _) => i16::try_from(*v).ok(), #[cfg(feature = "decimal")] - Scalar::Decimal128(v, _) => i16::try_from(*v).ok(), + Scalar::Decimal128(v, _, _) => i16::try_from(*v).ok(), Scalar::Null => None, Scalar::String32(s) => s.parse::().ok(), #[cfg(feature = "large_string")] @@ -1120,11 +1122,11 @@ impl Scalar { Scalar::Float32(v) => Some(*v as i32), Scalar::Float64(v) => Some(*v as i32), #[cfg(feature = "decimal")] - Scalar::Decimal32(v, _) => Some(*v), + Scalar::Decimal32(v, _, _) => Some(*v), #[cfg(feature = "decimal")] - Scalar::Decimal64(v, _) => i32::try_from(*v).ok(), + Scalar::Decimal64(v, _, _) => i32::try_from(*v).ok(), #[cfg(feature = "decimal")] - Scalar::Decimal128(v, _) => i32::try_from(*v).ok(), + Scalar::Decimal128(v, _, _) => i32::try_from(*v).ok(), Scalar::Null => None, Scalar::String32(s) => s.parse::().ok(), #[cfg(feature = "large_string")] @@ -1164,11 +1166,11 @@ impl Scalar { Scalar::Float32(v) => Some(*v as i64), Scalar::Float64(v) => Some(*v as i64), #[cfg(feature = "decimal")] - Scalar::Decimal32(v, _) => Some(*v as i64), + Scalar::Decimal32(v, _, _) => Some(*v as i64), #[cfg(feature = "decimal")] - Scalar::Decimal64(v, _) => Some(*v), + Scalar::Decimal64(v, _, _) => Some(*v), #[cfg(feature = "decimal")] - Scalar::Decimal128(v, _) => i64::try_from(*v).ok(), + Scalar::Decimal128(v, _, _) => i64::try_from(*v).ok(), Scalar::Null => None, Scalar::String32(s) => s.parse::().ok(), #[cfg(feature = "large_string")] @@ -1202,11 +1204,11 @@ impl Scalar { Scalar::Float32(v) => u8::try_from(*v as i32).ok(), Scalar::Float64(v) => u8::try_from(*v as i32).ok(), #[cfg(feature = "decimal")] - Scalar::Decimal32(v, _) => u8::try_from(*v).ok(), + Scalar::Decimal32(v, _, _) => u8::try_from(*v).ok(), #[cfg(feature = "decimal")] - Scalar::Decimal64(v, _) => u8::try_from(*v).ok(), + Scalar::Decimal64(v, _, _) => u8::try_from(*v).ok(), #[cfg(feature = "decimal")] - Scalar::Decimal128(v, _) => u8::try_from(*v).ok(), + Scalar::Decimal128(v, _, _) => u8::try_from(*v).ok(), Scalar::Null => None, Scalar::String32(s) => s.parse::().ok(), #[cfg(feature = "large_string")] @@ -1240,11 +1242,11 @@ impl Scalar { Scalar::Float32(v) => u16::try_from(*v as i32).ok(), Scalar::Float64(v) => u16::try_from(*v as i32).ok(), #[cfg(feature = "decimal")] - Scalar::Decimal32(v, _) => u16::try_from(*v).ok(), + Scalar::Decimal32(v, _, _) => u16::try_from(*v).ok(), #[cfg(feature = "decimal")] - Scalar::Decimal64(v, _) => u16::try_from(*v).ok(), + Scalar::Decimal64(v, _, _) => u16::try_from(*v).ok(), #[cfg(feature = "decimal")] - Scalar::Decimal128(v, _) => u16::try_from(*v).ok(), + Scalar::Decimal128(v, _, _) => u16::try_from(*v).ok(), Scalar::Null => None, Scalar::String32(s) => s.parse::().ok(), #[cfg(feature = "large_string")] @@ -1278,11 +1280,11 @@ impl Scalar { Scalar::Float32(v) => Some(*v as u32), Scalar::Float64(v) => Some(*v as u32), #[cfg(feature = "decimal")] - Scalar::Decimal32(v, _) => u32::try_from(*v).ok(), + Scalar::Decimal32(v, _, _) => u32::try_from(*v).ok(), #[cfg(feature = "decimal")] - Scalar::Decimal64(v, _) => u32::try_from(*v).ok(), + Scalar::Decimal64(v, _, _) => u32::try_from(*v).ok(), #[cfg(feature = "decimal")] - Scalar::Decimal128(v, _) => u32::try_from(*v).ok(), + Scalar::Decimal128(v, _, _) => u32::try_from(*v).ok(), Scalar::Null => None, Scalar::String32(s) => s.parse::().ok(), #[cfg(feature = "large_string")] @@ -1352,11 +1354,11 @@ impl Scalar { } } #[cfg(feature = "decimal")] - Scalar::Decimal32(v, _) => if *v >= 0 { Some(*v as u64) } else { None }, + Scalar::Decimal32(v, _, _) => if *v >= 0 { Some(*v as u64) } else { None }, #[cfg(feature = "decimal")] - Scalar::Decimal64(v, _) => if *v >= 0 { Some(*v as u64) } else { None }, + Scalar::Decimal64(v, _, _) => if *v >= 0 { Some(*v as u64) } else { None }, #[cfg(feature = "decimal")] - Scalar::Decimal128(v, _) => u64::try_from(*v).ok(), + Scalar::Decimal128(v, _, _) => u64::try_from(*v).ok(), Scalar::Null => None, Scalar::String32(s) => s.parse::().ok(), #[cfg(feature = "large_string")] @@ -1402,11 +1404,11 @@ impl Scalar { Scalar::Float32(v) => Some(*v), Scalar::Float64(v) => Some(*v as f32), #[cfg(feature = "decimal")] - Scalar::Decimal32(v, s) => Some(*v as f32 / 10f32.powi(*s as i32)), + Scalar::Decimal32(v, _, s) => Some(*v as f32 / 10f32.powi(*s as i32)), #[cfg(feature = "decimal")] - Scalar::Decimal64(v, s) => Some(*v as f32 / 10f32.powi(*s as i32)), + Scalar::Decimal64(v, _, s) => Some(*v as f32 / 10f32.powi(*s as i32)), #[cfg(feature = "decimal")] - Scalar::Decimal128(v, s) => Some(*v as f32 / 10f32.powi(*s as i32)), + Scalar::Decimal128(v, _, s) => Some(*v as f32 / 10f32.powi(*s as i32)), Scalar::Boolean(v) => Some(if *v { 1.0 } else { 0.0 }), Scalar::Null => None, Scalar::String32(s) => s.parse::().ok(), @@ -1440,11 +1442,11 @@ impl Scalar { Scalar::Float32(v) => Some(*v as f64), Scalar::Float64(v) => Some(*v), #[cfg(feature = "decimal")] - Scalar::Decimal32(v, s) => Some(*v as f64 / 10f64.powi(*s as i32)), + Scalar::Decimal32(v, _, s) => Some(*v as f64 / 10f64.powi(*s as i32)), #[cfg(feature = "decimal")] - Scalar::Decimal64(v, s) => Some(*v as f64 / 10f64.powi(*s as i32)), + Scalar::Decimal64(v, _, s) => Some(*v as f64 / 10f64.powi(*s as i32)), #[cfg(feature = "decimal")] - Scalar::Decimal128(v, s) => Some(*v as f64 / 10f64.powi(*s as i32)), + Scalar::Decimal128(v, _, s) => Some(*v as f64 / 10f64.powi(*s as i32)), Scalar::Boolean(v) => Some(if *v { 1.0 } else { 0.0 }), Scalar::Null => None, Scalar::String32(s) => s.parse::().ok(), @@ -1482,11 +1484,11 @@ impl Scalar { Scalar::Float32(v) => Some(v.to_string()), Scalar::Float64(v) => Some(v.to_string()), #[cfg(feature = "decimal")] - Scalar::Decimal32(v, s) => Some(format_decimal_scalar(*v as i128, *s)), + Scalar::Decimal32(v, _, s) => Some(format_decimal_scalar(*v as i128, *s)), #[cfg(feature = "decimal")] - Scalar::Decimal64(v, s) => Some(format_decimal_scalar(*v as i128, *s)), + Scalar::Decimal64(v, _, s) => Some(format_decimal_scalar(*v as i128, *s)), #[cfg(feature = "decimal")] - Scalar::Decimal128(v, s) => Some(format_decimal_scalar(*v, *s)), + Scalar::Decimal128(v, _, s) => Some(format_decimal_scalar(*v, *s)), Scalar::Null => None, #[cfg(feature = "datetime")] Scalar::Datetime32(v) => Some(v.to_string()), @@ -1548,7 +1550,7 @@ impl Scalar { } } #[cfg(feature = "decimal")] - Scalar::Decimal32(_, _) | Scalar::Decimal64(_, _) | Scalar::Decimal128(_, _) => None, + Scalar::Decimal32(_, _, _) | Scalar::Decimal64(_, _, _) | Scalar::Decimal128(_, _, _) => None, Scalar::String32(s) => s.parse::().ok(), #[cfg(feature = "large_string")] Scalar::String64(s) => s.parse::().ok(), @@ -1628,7 +1630,7 @@ impl Scalar { } } #[cfg(feature = "decimal")] - Scalar::Decimal32(_, _) | Scalar::Decimal64(_, _) | Scalar::Decimal128(_, _) => None, + Scalar::Decimal32(_, _, _) | Scalar::Decimal64(_, _, _) | Scalar::Decimal128(_, _, _) => None, Scalar::String32(s) => s.parse::().ok(), #[cfg(feature = "large_string")] Scalar::String64(s) => s.parse::().ok(), @@ -1677,11 +1679,11 @@ impl Scalar { Scalar::Float32(_) => ArrowType::Float32, Scalar::Float64(_) => ArrowType::Float64, #[cfg(feature = "decimal")] - Scalar::Decimal32(_, s) => ArrowType::Decimal32(0, *s), + Scalar::Decimal32(_, p, s) => ArrowType::Decimal32(*p, *s), #[cfg(feature = "decimal")] - Scalar::Decimal64(_, s) => ArrowType::Decimal64(0, *s), + Scalar::Decimal64(_, p, s) => ArrowType::Decimal64(*p, *s), #[cfg(feature = "decimal")] - Scalar::Decimal128(_, s) => ArrowType::Decimal128(0, *s), + Scalar::Decimal128(_, p, s) => ArrowType::Decimal128(*p, *s), Scalar::String32(_) => ArrowType::String, #[cfg(feature = "large_string")] Scalar::String64(_) => ArrowType::LargeString, @@ -1772,24 +1774,24 @@ impl Scalar { Array::from_float64(arr) } #[cfg(feature = "decimal")] - Scalar::Decimal32(v, s) => { - let mut arr = crate::DecimalArray::::with_capacity(len, false, 0, s); + Scalar::Decimal32(v, p, s) => { + let mut arr = crate::DecimalArray::::with_capacity(len, false, p, s); for _ in 0..len { arr.push(v); } Array::NumericArray(crate::NumericArray::Decimal32(std::sync::Arc::new(arr))) } #[cfg(feature = "decimal")] - Scalar::Decimal64(v, s) => { - let mut arr = crate::DecimalArray::::with_capacity(len, false, 0, s); + Scalar::Decimal64(v, p, s) => { + let mut arr = crate::DecimalArray::::with_capacity(len, false, p, s); for _ in 0..len { arr.push(v); } Array::NumericArray(crate::NumericArray::Decimal64(std::sync::Arc::new(arr))) } #[cfg(feature = "decimal")] - Scalar::Decimal128(v, s) => { - let mut arr = crate::DecimalArray::::with_capacity(len, false, 0, s); + Scalar::Decimal128(v, p, s) => { + let mut arr = crate::DecimalArray::::with_capacity(len, false, p, s); for _ in 0..len { arr.push(v); } @@ -1873,11 +1875,13 @@ impl PartialEq for Scalar { (Float32(a), Float32(b)) => if a.is_nan() { b.is_nan() } else { a == b }, (Float64(a), Float64(b)) => if a.is_nan() { b.is_nan() } else { a == b }, #[cfg(feature = "decimal")] - (Decimal32(a, sa), Decimal32(b, sb)) => a == b && sa == sb, + // Decimal equality covers value, precision and scale, matching + // `ArrowType` equality for the column type. + (Decimal32(a, pa, sa), Decimal32(b, pb, sb)) => a == b && pa == pb && sa == sb, #[cfg(feature = "decimal")] - (Decimal64(a, sa), Decimal64(b, sb)) => a == b && sa == sb, + (Decimal64(a, pa, sa), Decimal64(b, pb, sb)) => a == b && pa == pb && sa == sb, #[cfg(feature = "decimal")] - (Decimal128(a, sa), Decimal128(b, sb)) => a == b && sa == sb, + (Decimal128(a, pa, sa), Decimal128(b, pb, sb)) => a == b && pa == pb && sa == sb, (String32(a), String32(b)) => a == b, #[cfg(feature = "large_string")] (String64(a), String64(b)) => a == b, @@ -1927,11 +1931,11 @@ impl std::hash::Hash for Scalar { bits.hash(state); } #[cfg(feature = "decimal")] - Scalar::Decimal32(v, s) => { v.hash(state); s.hash(state); } + Scalar::Decimal32(v, p, s) => { v.hash(state); p.hash(state); s.hash(state); } #[cfg(feature = "decimal")] - Scalar::Decimal64(v, s) => { v.hash(state); s.hash(state); } + Scalar::Decimal64(v, p, s) => { v.hash(state); p.hash(state); s.hash(state); } #[cfg(feature = "decimal")] - Scalar::Decimal128(v, s) => { v.hash(state); s.hash(state); } + Scalar::Decimal128(v, p, s) => { v.hash(state); p.hash(state); s.hash(state); } Scalar::String32(v) => v.hash(state), #[cfg(feature = "large_string")] Scalar::String64(v) => v.hash(state), @@ -2052,17 +2056,17 @@ impl Add for Scalar { // Decimal promotes to f64 #[cfg(feature = "decimal")] - (Decimal32(a, s), b) => Float64(a as f64 / 10f64.powi(s as i32) + b.f64()), + (Decimal32(a, _, s), b) => Float64(a as f64 / 10f64.powi(s as i32) + b.f64()), #[cfg(feature = "decimal")] - (a, Decimal32(b, s)) => Float64(a.f64() + b as f64 / 10f64.powi(s as i32)), + (a, Decimal32(b, _, s)) => Float64(a.f64() + b as f64 / 10f64.powi(s as i32)), #[cfg(feature = "decimal")] - (Decimal64(a, s), b) => Float64(a as f64 / 10f64.powi(s as i32) + b.f64()), + (Decimal64(a, _, s), b) => Float64(a as f64 / 10f64.powi(s as i32) + b.f64()), #[cfg(feature = "decimal")] - (a, Decimal64(b, s)) => Float64(a.f64() + b as f64 / 10f64.powi(s as i32)), + (a, Decimal64(b, _, s)) => Float64(a.f64() + b as f64 / 10f64.powi(s as i32)), #[cfg(feature = "decimal")] - (Decimal128(a, s), b) => Float64(a as f64 / 10f64.powi(s as i32) + b.f64()), + (Decimal128(a, _, s), b) => Float64(a as f64 / 10f64.powi(s as i32) + b.f64()), #[cfg(feature = "decimal")] - (a, Decimal128(b, s)) => Float64(a.f64() + b as f64 / 10f64.powi(s as i32)), + (a, Decimal128(b, _, s)) => Float64(a.f64() + b as f64 / 10f64.powi(s as i32)), // Float promotion (Float64(a), b) => Float64(a + b.f64()), @@ -2148,17 +2152,17 @@ impl Sub for Scalar { (Null, _) | (_, Null) => Null, #[cfg(feature = "decimal")] - (Decimal32(a, s), b) => Float64(a as f64 / 10f64.powi(s as i32) - b.f64()), + (Decimal32(a, _, s), b) => Float64(a as f64 / 10f64.powi(s as i32) - b.f64()), #[cfg(feature = "decimal")] - (a, Decimal32(b, s)) => Float64(a.f64() - b as f64 / 10f64.powi(s as i32)), + (a, Decimal32(b, _, s)) => Float64(a.f64() - b as f64 / 10f64.powi(s as i32)), #[cfg(feature = "decimal")] - (Decimal64(a, s), b) => Float64(a as f64 / 10f64.powi(s as i32) - b.f64()), + (Decimal64(a, _, s), b) => Float64(a as f64 / 10f64.powi(s as i32) - b.f64()), #[cfg(feature = "decimal")] - (a, Decimal64(b, s)) => Float64(a.f64() - b as f64 / 10f64.powi(s as i32)), + (a, Decimal64(b, _, s)) => Float64(a.f64() - b as f64 / 10f64.powi(s as i32)), #[cfg(feature = "decimal")] - (Decimal128(a, s), b) => Float64(a as f64 / 10f64.powi(s as i32) - b.f64()), + (Decimal128(a, _, s), b) => Float64(a as f64 / 10f64.powi(s as i32) - b.f64()), #[cfg(feature = "decimal")] - (a, Decimal128(b, s)) => Float64(a.f64() - b as f64 / 10f64.powi(s as i32)), + (a, Decimal128(b, _, s)) => Float64(a.f64() - b as f64 / 10f64.powi(s as i32)), (Float64(a), b) => Float64(a - b.f64()), (a, Float64(b)) => Float64(a.f64() - b), @@ -2232,17 +2236,17 @@ impl Mul for Scalar { (Null, _) | (_, Null) => Null, #[cfg(feature = "decimal")] - (Decimal32(a, s), b) => Float64(a as f64 / 10f64.powi(s as i32) * b.f64()), + (Decimal32(a, _, s), b) => Float64(a as f64 / 10f64.powi(s as i32) * b.f64()), #[cfg(feature = "decimal")] - (a, Decimal32(b, s)) => Float64(a.f64() * (b as f64 / 10f64.powi(s as i32))), + (a, Decimal32(b, _, s)) => Float64(a.f64() * (b as f64 / 10f64.powi(s as i32))), #[cfg(feature = "decimal")] - (Decimal64(a, s), b) => Float64(a as f64 / 10f64.powi(s as i32) * b.f64()), + (Decimal64(a, _, s), b) => Float64(a as f64 / 10f64.powi(s as i32) * b.f64()), #[cfg(feature = "decimal")] - (a, Decimal64(b, s)) => Float64(a.f64() * (b as f64 / 10f64.powi(s as i32))), + (a, Decimal64(b, _, s)) => Float64(a.f64() * (b as f64 / 10f64.powi(s as i32))), #[cfg(feature = "decimal")] - (Decimal128(a, s), b) => Float64(a as f64 / 10f64.powi(s as i32) * b.f64()), + (Decimal128(a, _, s), b) => Float64(a as f64 / 10f64.powi(s as i32) * b.f64()), #[cfg(feature = "decimal")] - (a, Decimal128(b, s)) => Float64(a.f64() * (b as f64 / 10f64.powi(s as i32))), + (a, Decimal128(b, _, s)) => Float64(a.f64() * (b as f64 / 10f64.powi(s as i32))), (Float64(a), b) => Float64(a * b.f64()), (a, Float64(b)) => Float64(a.f64() * b), @@ -2343,17 +2347,17 @@ impl Pow for Scalar { (Null, _) | (_, Null) => Null, #[cfg(feature = "decimal")] - (Decimal32(a, s), b) => Float64((a as f64 / 10f64.powi(s as i32)).powf(b.f64())), + (Decimal32(a, _, s), b) => Float64((a as f64 / 10f64.powi(s as i32)).powf(b.f64())), #[cfg(feature = "decimal")] - (a, Decimal32(b, s)) => Float64(a.f64().powf(b as f64 / 10f64.powi(s as i32))), + (a, Decimal32(b, _, s)) => Float64(a.f64().powf(b as f64 / 10f64.powi(s as i32))), #[cfg(feature = "decimal")] - (Decimal64(a, s), b) => Float64((a as f64 / 10f64.powi(s as i32)).powf(b.f64())), + (Decimal64(a, _, s), b) => Float64((a as f64 / 10f64.powi(s as i32)).powf(b.f64())), #[cfg(feature = "decimal")] - (a, Decimal64(b, s)) => Float64(a.f64().powf(b as f64 / 10f64.powi(s as i32))), + (a, Decimal64(b, _, s)) => Float64(a.f64().powf(b as f64 / 10f64.powi(s as i32))), #[cfg(feature = "decimal")] - (Decimal128(a, s), b) => Float64((a as f64 / 10f64.powi(s as i32)).powf(b.f64())), + (Decimal128(a, _, s), b) => Float64((a as f64 / 10f64.powi(s as i32)).powf(b.f64())), #[cfg(feature = "decimal")] - (a, Decimal128(b, s)) => Float64(a.f64().powf(b as f64 / 10f64.powi(s as i32))), + (a, Decimal128(b, _, s)) => Float64(a.f64().powf(b as f64 / 10f64.powi(s as i32))), #[cfg(feature = "datetime")] (Interval, _) => panic!("Cannot exponentiate Interval"), @@ -2528,6 +2532,28 @@ mod tests { ArrowType::Interval(IntervalUnit::MonthDaysNs) ); } + + // Decimal scalars carry precision and scale, so the reported type is + // the exact column type. + #[cfg(feature = "decimal")] + { + assert_eq!(Scalar::Decimal32(1, 9, 2).arrow_type(), ArrowType::Decimal32(9, 2)); + assert_eq!(Scalar::Decimal64(1, 18, 4).arrow_type(), ArrowType::Decimal64(18, 4)); + assert_eq!( + Scalar::Decimal128(1, 38, 10).arrow_type(), + ArrowType::Decimal128(38, 10) + ); + } + } + + #[cfg(feature = "decimal")] + #[test] + fn decimal_equality_covers_value_precision_and_scale() { + let a = Scalar::Decimal64(10050, 18, 4); + assert_eq!(a, Scalar::Decimal64(10050, 18, 4)); + assert_ne!(a, Scalar::Decimal64(10050, 12, 4)); + assert_ne!(a, Scalar::Decimal64(10050, 18, 2)); + assert_ne!(a, Scalar::Decimal64(10051, 18, 4)); } #[test] diff --git a/src/enums/value/impls.rs b/src/enums/value/impls.rs index 25af0e9..764d3cf 100644 --- a/src/enums/value/impls.rs +++ b/src/enums/value/impls.rs @@ -615,11 +615,11 @@ fn scalar_variant_name(scalar: &crate::Scalar) -> &'static str { #[cfg(feature = "datetime")] Interval => "Interval", #[cfg(feature = "decimal")] - Decimal32(_, _) => "Decimal32", + Decimal32(_, _, _) => "Decimal32", #[cfg(feature = "decimal")] - Decimal64(_, _) => "Decimal64", + Decimal64(_, _, _) => "Decimal64", #[cfg(feature = "decimal")] - Decimal128(_, _) => "Decimal128", + Decimal128(_, _, _) => "Decimal128", } } diff --git a/src/kernels/arithmetic/decimal.rs b/src/kernels/arithmetic/decimal.rs index 26bd999..82746d5 100644 --- a/src/kernels/arithmetic/decimal.rs +++ b/src/kernels/arithmetic/decimal.rs @@ -26,7 +26,7 @@ use crate::traits::type_unions::Integer; use crate::{Bitmask, DecimalArray, MaskedArray, Vec64}; /// Maximum precision per backing integer width. -fn max_precision() -> u8 { +pub(crate) fn max_precision() -> u8 { use std::any::TypeId; let tid = TypeId::of::(); if tid == TypeId::of::() { diff --git a/src/kernels/broadcast/array.rs b/src/kernels/broadcast/array.rs index c0eb8f9..e370e02 100644 --- a/src/kernels/broadcast/array.rs +++ b/src/kernels/broadcast/array.rs @@ -173,16 +173,16 @@ pub fn broadcast_array_to_scalar( Array::from_datetime_i64(DatetimeArray::from_slice(&[*val], None)) } #[cfg(feature = "decimal")] - Scalar::Decimal32(val, s) => { - Array::from_decimal32(crate::DecimalArray::from_slice(&[*val], 0, *s)) + Scalar::Decimal32(val, p, s) => { + Array::from_decimal32(crate::DecimalArray::from_slice(&[*val], *p, *s)) } #[cfg(feature = "decimal")] - Scalar::Decimal64(val, s) => { - Array::from_decimal64(crate::DecimalArray::from_slice(&[*val], 0, *s)) + Scalar::Decimal64(val, p, s) => { + Array::from_decimal64(crate::DecimalArray::from_slice(&[*val], *p, *s)) } #[cfg(feature = "decimal")] - Scalar::Decimal128(val, s) => { - Array::from_decimal128(crate::DecimalArray::from_slice(&[*val], 0, *s)) + Scalar::Decimal128(val, p, s) => { + Array::from_decimal128(crate::DecimalArray::from_slice(&[*val], *p, *s)) } Scalar::Null => Array::Null, #[cfg(feature = "datetime")] diff --git a/src/kernels/broadcast/scalar.rs b/src/kernels/broadcast/scalar.rs index 9f76d86..549fe72 100644 --- a/src/kernels/broadcast/scalar.rs +++ b/src/kernels/broadcast/scalar.rs @@ -203,16 +203,16 @@ pub fn broadcast_scalar_to_array( Array::from_datetime_i64(DatetimeArray::from_slice(&[*val], None)) } #[cfg(feature = "decimal")] - Scalar::Decimal32(val, s) => { - Array::from_decimal32(crate::DecimalArray::from_slice(&[*val], 0, *s)) + Scalar::Decimal32(val, p, s) => { + Array::from_decimal32(crate::DecimalArray::from_slice(&[*val], *p, *s)) } #[cfg(feature = "decimal")] - Scalar::Decimal64(val, s) => { - Array::from_decimal64(crate::DecimalArray::from_slice(&[*val], 0, *s)) + Scalar::Decimal64(val, p, s) => { + Array::from_decimal64(crate::DecimalArray::from_slice(&[*val], *p, *s)) } #[cfg(feature = "decimal")] - Scalar::Decimal128(val, s) => { - Array::from_decimal128(crate::DecimalArray::from_slice(&[*val], 0, *s)) + Scalar::Decimal128(val, p, s) => { + Array::from_decimal128(crate::DecimalArray::from_slice(&[*val], *p, *s)) } Scalar::Null => Array::Null, #[cfg(feature = "datetime")] @@ -647,9 +647,9 @@ pub fn broadcast_scalar_to_text_arrayview( }); } #[cfg(feature = "decimal")] - (Scalar::Decimal32(_, _), _) - | (Scalar::Decimal64(_, _), _) - | (Scalar::Decimal128(_, _), _) => { + (Scalar::Decimal32(_, _, _), _) + | (Scalar::Decimal64(_, _, _), _) + | (Scalar::Decimal128(_, _, _), _) => { return Err(MinarrowError::NotImplemented { feature: "Decimal scalar with TextArrayView".to_string(), }); @@ -763,9 +763,9 @@ pub fn broadcast_text_arrayview_to_scalar( }); } #[cfg(feature = "decimal")] - (_, Scalar::Decimal32(_, _)) - | (_, Scalar::Decimal64(_, _)) - | (_, Scalar::Decimal128(_, _)) => { + (_, Scalar::Decimal32(_, _, _)) + | (_, Scalar::Decimal64(_, _, _)) + | (_, Scalar::Decimal128(_, _, _)) => { return Err(MinarrowError::NotImplemented { feature: "Decimal scalar with TextArrayView".to_string(), }); @@ -822,16 +822,16 @@ pub fn broadcast_scalar_to_fieldarray( Array::from_datetime_i64(DatetimeArray::from_slice(&[*val], None)) } #[cfg(feature = "decimal")] - Scalar::Decimal32(val, s) => { - Array::from_decimal32(crate::DecimalArray::from_slice(&[*val], 0, *s)) + Scalar::Decimal32(val, p, s) => { + Array::from_decimal32(crate::DecimalArray::from_slice(&[*val], *p, *s)) } #[cfg(feature = "decimal")] - Scalar::Decimal64(val, s) => { - Array::from_decimal64(crate::DecimalArray::from_slice(&[*val], 0, *s)) + Scalar::Decimal64(val, p, s) => { + Array::from_decimal64(crate::DecimalArray::from_slice(&[*val], *p, *s)) } #[cfg(feature = "decimal")] - Scalar::Decimal128(val, s) => { - Array::from_decimal128(crate::DecimalArray::from_slice(&[*val], 0, *s)) + Scalar::Decimal128(val, p, s) => { + Array::from_decimal128(crate::DecimalArray::from_slice(&[*val], *p, *s)) } Scalar::Null => Array::Null, #[cfg(feature = "datetime")] @@ -879,16 +879,16 @@ pub fn broadcast_fieldarray_to_scalar( Array::from_datetime_i64(DatetimeArray::from_slice(&[*val], None)) } #[cfg(feature = "decimal")] - Scalar::Decimal32(val, s) => { - Array::from_decimal32(crate::DecimalArray::from_slice(&[*val], 0, *s)) + Scalar::Decimal32(val, p, s) => { + Array::from_decimal32(crate::DecimalArray::from_slice(&[*val], *p, *s)) } #[cfg(feature = "decimal")] - Scalar::Decimal64(val, s) => { - Array::from_decimal64(crate::DecimalArray::from_slice(&[*val], 0, *s)) + Scalar::Decimal64(val, p, s) => { + Array::from_decimal64(crate::DecimalArray::from_slice(&[*val], *p, *s)) } #[cfg(feature = "decimal")] - Scalar::Decimal128(val, s) => { - Array::from_decimal128(crate::DecimalArray::from_slice(&[*val], 0, *s)) + Scalar::Decimal128(val, p, s) => { + Array::from_decimal128(crate::DecimalArray::from_slice(&[*val], *p, *s)) } Scalar::Null => Array::Null, #[cfg(feature = "datetime")] @@ -947,7 +947,7 @@ pub fn broadcast_scalar_to_temporal_arrayview( }); } #[cfg(feature = "decimal")] - Scalar::Decimal32(_, _) | Scalar::Decimal64(_, _) | Scalar::Decimal128(_, _) => { + Scalar::Decimal32(_, _, _) | Scalar::Decimal64(_, _, _) | Scalar::Decimal128(_, _, _) => { return Err(MinarrowError::NotImplemented { feature: "Decimal scalar with TemporalArrayView".to_string(), }); @@ -1005,7 +1005,7 @@ pub fn broadcast_temporal_arrayview_to_scalar( }); } #[cfg(feature = "decimal")] - Scalar::Decimal32(_, _) | Scalar::Decimal64(_, _) | Scalar::Decimal128(_, _) => { + Scalar::Decimal32(_, _, _) | Scalar::Decimal64(_, _, _) | Scalar::Decimal128(_, _, _) => { return Err(MinarrowError::NotImplemented { feature: "Decimal scalar with TemporalArrayView".to_string(), }); diff --git a/src/kernels/routing/arithmetic.rs b/src/kernels/routing/arithmetic.rs index e9396f0..1e464a0 100644 --- a/src/kernels/routing/arithmetic.rs +++ b/src/kernels/routing/arithmetic.rs @@ -29,7 +29,7 @@ use crate::kernels::arithmetic::{ string_ops::apply_str_str, }; #[cfg(feature = "decimal")] -use crate::kernels::arithmetic::decimal::{decimal_binary, integer_to_decimal}; +use crate::kernels::arithmetic::decimal::{decimal_binary, integer_to_decimal, max_precision}; use crate::enums::{error::KernelError, operators::ArithmeticOperator}; @@ -193,75 +193,80 @@ pub fn scalar_arithmetic( #[cfg(feature = "large_string")] (Scalar::String64(l), Scalar::String32(r), Add) => Scalar::String64(format!("{}{}", l, r)), - // Decimal scalar operations at matching width and scale + // Decimal scalar operations at matching width and scale. + // + // ## Behaviour + // - add and subtract widen the larger operand precision by one digit + // - multiply sums the operand precisions + // - both are capped at the width maximum. #[cfg(feature = "decimal")] - (Scalar::Decimal32(l, ls), Scalar::Decimal32(r, rs), Add) if ls == rs => { + (Scalar::Decimal32(l, lp, ls), Scalar::Decimal32(r, rp, rs), Add) if ls == rs => { Scalar::Decimal32(l.checked_add(r).ok_or_else(|| MinarrowError::KernelError( Some("Decimal32 overflow in addition".to_string()), - ))?, ls) + ))?, (lp.max(rp) + 1).min(max_precision::()), ls) } #[cfg(feature = "decimal")] - (Scalar::Decimal32(l, ls), Scalar::Decimal32(r, rs), Subtract) if ls == rs => { + (Scalar::Decimal32(l, lp, ls), Scalar::Decimal32(r, rp, rs), Subtract) if ls == rs => { Scalar::Decimal32(l.checked_sub(r).ok_or_else(|| MinarrowError::KernelError( Some("Decimal32 overflow in subtraction".to_string()), - ))?, ls) + ))?, (lp.max(rp) + 1).min(max_precision::()), ls) } #[cfg(feature = "decimal")] - (Scalar::Decimal32(l, ls), Scalar::Decimal32(r, _rs), Multiply) => { + (Scalar::Decimal32(l, lp, ls), Scalar::Decimal32(r, rp, rs), Multiply) => { Scalar::Decimal32(l.checked_mul(r).ok_or_else(|| MinarrowError::KernelError( Some("Decimal32 overflow in multiplication".to_string()), - ))?, ls + _rs) + ))?, lp.saturating_add(rp).min(max_precision::()), ls + rs) } #[cfg(feature = "decimal")] - (Scalar::Decimal64(l, ls), Scalar::Decimal64(r, rs), Add) if ls == rs => { + (Scalar::Decimal64(l, lp, ls), Scalar::Decimal64(r, rp, rs), Add) if ls == rs => { Scalar::Decimal64(l.checked_add(r).ok_or_else(|| MinarrowError::KernelError( Some("Decimal64 overflow in addition".to_string()), - ))?, ls) + ))?, (lp.max(rp) + 1).min(max_precision::()), ls) } #[cfg(feature = "decimal")] - (Scalar::Decimal64(l, ls), Scalar::Decimal64(r, rs), Subtract) if ls == rs => { + (Scalar::Decimal64(l, lp, ls), Scalar::Decimal64(r, rp, rs), Subtract) if ls == rs => { Scalar::Decimal64(l.checked_sub(r).ok_or_else(|| MinarrowError::KernelError( Some("Decimal64 overflow in subtraction".to_string()), - ))?, ls) + ))?, (lp.max(rp) + 1).min(max_precision::()), ls) } #[cfg(feature = "decimal")] - (Scalar::Decimal64(l, ls), Scalar::Decimal64(r, _rs), Multiply) => { + (Scalar::Decimal64(l, lp, ls), Scalar::Decimal64(r, rp, rs), Multiply) => { Scalar::Decimal64(l.checked_mul(r).ok_or_else(|| MinarrowError::KernelError( Some("Decimal64 overflow in multiplication".to_string()), - ))?, ls + _rs) + ))?, lp.saturating_add(rp).min(max_precision::()), ls + rs) } #[cfg(feature = "decimal")] - (Scalar::Decimal128(l, ls), Scalar::Decimal128(r, rs), Add) if ls == rs => { + (Scalar::Decimal128(l, lp, ls), Scalar::Decimal128(r, rp, rs), Add) if ls == rs => { Scalar::Decimal128(l.checked_add(r).ok_or_else(|| MinarrowError::KernelError( Some("Decimal128 overflow in addition".to_string()), - ))?, ls) + ))?, (lp.max(rp) + 1).min(max_precision::()), ls) } #[cfg(feature = "decimal")] - (Scalar::Decimal128(l, ls), Scalar::Decimal128(r, rs), Subtract) if ls == rs => { + (Scalar::Decimal128(l, lp, ls), Scalar::Decimal128(r, rp, rs), Subtract) if ls == rs => { Scalar::Decimal128(l.checked_sub(r).ok_or_else(|| MinarrowError::KernelError( Some("Decimal128 overflow in subtraction".to_string()), - ))?, ls) + ))?, (lp.max(rp) + 1).min(max_precision::()), ls) } #[cfg(feature = "decimal")] - (Scalar::Decimal128(l, ls), Scalar::Decimal128(r, _rs), Multiply) => { + (Scalar::Decimal128(l, lp, ls), Scalar::Decimal128(r, rp, rs), Multiply) => { Scalar::Decimal128(l.checked_mul(r).ok_or_else(|| MinarrowError::KernelError( Some("Decimal128 overflow in multiplication".to_string()), - ))?, ls + _rs) + ))?, lp.saturating_add(rp).min(max_precision::()), ls + rs) } // Decimal + Float -> Float64 scalar promotion #[cfg(feature = "decimal")] - (Scalar::Decimal32(l, s), Scalar::Float64(r), op) | (Scalar::Float64(r), Scalar::Decimal32(l, s), op) => { + (Scalar::Decimal32(l, _, s), Scalar::Float64(r), op) | (Scalar::Float64(r), Scalar::Decimal32(l, _, s), op) => { let l_f64 = l as f64 / 10f64.powi(s as i32); return scalar_arithmetic(Scalar::Float64(l_f64), Scalar::Float64(r), op); } #[cfg(feature = "decimal")] - (Scalar::Decimal64(l, s), Scalar::Float64(r), op) | (Scalar::Float64(r), Scalar::Decimal64(l, s), op) => { + (Scalar::Decimal64(l, _, s), Scalar::Float64(r), op) | (Scalar::Float64(r), Scalar::Decimal64(l, _, s), op) => { let l_f64 = l as f64 / 10f64.powi(s as i32); return scalar_arithmetic(Scalar::Float64(l_f64), Scalar::Float64(r), op); } #[cfg(feature = "decimal")] - (Scalar::Decimal128(l, s), Scalar::Float64(r), op) | (Scalar::Float64(r), Scalar::Decimal128(l, s), op) => { + (Scalar::Decimal128(l, _, s), Scalar::Float64(r), op) | (Scalar::Float64(r), Scalar::Decimal128(l, _, s), op) => { use num_traits::ToPrimitive; let l_f64 = l.to_f64().unwrap() / 10f64.powi(s as i32); return scalar_arithmetic(Scalar::Float64(l_f64), Scalar::Float64(r), op); diff --git a/src/structs/field_array.rs b/src/structs/field_array.rs index 60e9f4f..f711342 100644 --- a/src/structs/field_array.rs +++ b/src/structs/field_array.rs @@ -1946,6 +1946,55 @@ mod concat_tests { } } + #[cfg(all(feature = "decimal", feature = "scalar_type"))] + #[test] + fn test_field_array_concat_decimal_rebuilt_from_scalars() { + use crate::{DecimalArray, Scalar}; + + let typed = FieldArray::from_arr( + "price", + Array::from_decimal64(DecimalArray::::from_slice(&[10050, 20075], 18, 4)), + ); + let rebuilt = FieldArray::from_arr( + "price", + Array::from_scalars(&[Scalar::Decimal64(30010, 18, 4)]), + ); + assert_eq!(rebuilt.field.dtype, ArrowType::Decimal64(18, 4)); + + let result = typed.concat(rebuilt).unwrap(); + assert_eq!(result.len(), 3); + assert_eq!(result.field.dtype, ArrowType::Decimal64(18, 4)); + match &result.array { + Array::NumericArray(NumericArray::Decimal64(arr)) => { + assert_eq!(arr.precision, 18); + assert_eq!(arr.get(2), Some(30010)); + } + other => panic!("Expected Decimal64 array, got {:?}", other), + } + } + + #[cfg(feature = "decimal")] + #[test] + fn test_field_array_concat_differing_decimal_precision_is_a_mismatch() { + use crate::DecimalArray; + + let fa1 = FieldArray::from_arr( + "price", + Array::from_decimal64(DecimalArray::::from_slice(&[10050], 18, 4)), + ); + let fa2 = FieldArray::from_arr( + "price", + Array::from_decimal64(DecimalArray::::from_slice(&[20075], 12, 4)), + ); + + let result = fa1.concat(fa2); + if let Err(MinarrowError::IncompatibleTypeError { message, .. }) = result { + assert!(message.unwrap().contains("dtype mismatch")); + } else { + panic!("Expected IncompatibleTypeError"); + } + } + #[test] fn test_field_array_concat_nullable_mismatch() { let arr1 = IntegerArray::::from_slice(&[1, 2]); diff --git a/src/structs/table.rs b/src/structs/table.rs index 6df7cd8..a6d2960 100644 --- a/src/structs/table.rs +++ b/src/structs/table.rs @@ -1764,6 +1764,137 @@ mod tests { assert!(result.is_err()); } + // Decimal columns rebuilt from scalars carry the full column type + + #[cfg(all(feature = "decimal", feature = "scalar_type"))] + #[test] + fn test_table_concat_decimal_column_rebuilt_from_scalars() { + use crate::ffi::arrow_dtype::ArrowType; + use crate::{DecimalArray, MaskedArray, Scalar}; + + let mut target = Table::new_empty(); + target.add_col(fa_i64!("id", 1, 2)); + target.add_col(FieldArray::from_arr( + "price", + Array::from_decimal64(DecimalArray::::from_slice(&[10050, 20075], 18, 4)), + )); + + let mut added = Table::new_empty(); + added.add_col(fa_i64!("id", 3)); + added.add_col(FieldArray::from_arr( + "price", + Array::from_scalars(&[Scalar::Decimal64(30010, 18, 4)]), + )); + assert_eq!(added.cols[1].field.dtype, ArrowType::Decimal64(18, 4)); + + let combined = target.concat(added).unwrap(); + assert_eq!(combined.n_rows(), 3); + assert_eq!(combined.cols[1].field.dtype, ArrowType::Decimal64(18, 4)); + match &combined.cols[1].array { + Array::NumericArray(NumericArray::Decimal64(arr)) => { + assert_eq!(arr.precision, 18); + assert_eq!(arr.scale, 4); + assert_eq!(arr.get(2), Some(30010)); + } + other => panic!("Expected Decimal64 array, got {:?}", other), + } + } + + #[cfg(all(feature = "decimal", feature = "scalar_type"))] + #[test] + fn test_table_concat_scalar_rebuilt_decimal_column_first_operand() { + use crate::ffi::arrow_dtype::ArrowType; + use crate::{DecimalArray, Scalar}; + + let mut rebuilt = Table::new_empty(); + rebuilt.add_col(FieldArray::from_arr( + "price", + Array::from_scalars(&[Scalar::Decimal64(30010, 18, 4)]), + )); + + let mut typed = Table::new_empty(); + typed.add_col(FieldArray::from_arr( + "price", + Array::from_decimal64(DecimalArray::::from_slice(&[10050], 18, 4)), + )); + + let combined = rebuilt.concat(typed).unwrap(); + assert_eq!(combined.n_rows(), 2); + assert_eq!(combined.cols[0].field.dtype, ArrowType::Decimal64(18, 4)); + match &combined.cols[0].array { + Array::NumericArray(NumericArray::Decimal64(arr)) => assert_eq!(arr.precision, 18), + other => panic!("Expected Decimal64 array, got {:?}", other), + } + } + + #[cfg(feature = "decimal")] + #[test] + fn test_table_concat_differing_decimal_precision_is_a_type_mismatch() { + use crate::DecimalArray; + + let mut t1 = Table::new_empty(); + t1.add_col(FieldArray::from_arr( + "price", + Array::from_decimal64(DecimalArray::::from_slice(&[10050], 18, 4)), + )); + + let mut t2 = Table::new_empty(); + t2.add_col(FieldArray::from_arr( + "price", + Array::from_decimal64(DecimalArray::::from_slice(&[20075], 12, 4)), + )); + + let err = t1.concat(t2).unwrap_err(); + assert!( + format!("{}", err).contains("type mismatch"), + "Expected type mismatch error, got: {}", + err + ); + } + + #[cfg(all(feature = "decimal", feature = "scalar_type"))] + #[test] + fn test_table_insert_rows_decimal_rows_rebuilt_from_scalars() { + use crate::ffi::arrow_dtype::ArrowType; + use crate::{DecimalArray, Scalar}; + + let mut typed = Table::new_empty(); + typed.add_col(FieldArray::from_arr( + "price", + Array::from_decimal64(DecimalArray::::from_slice(&[10050], 18, 4)), + )); + + let mut rebuilt = Table::new_empty(); + rebuilt.add_col(FieldArray::from_arr( + "price", + Array::from_scalars(&[Scalar::Decimal64(20075, 18, 4)]), + )); + + typed.insert_rows(1, &rebuilt).unwrap(); + assert_eq!(typed.n_rows(), 2); + assert_eq!(typed.cols[0].field.dtype, ArrowType::Decimal64(18, 4)); + } + + #[cfg(all(feature = "decimal", feature = "scalar_type"))] + #[test] + fn test_table_insert_rows_decimal_precision_mismatch_is_a_type_mismatch() { + use crate::{DecimalArray, Scalar}; + + let mut typed = Table::new_empty(); + typed.add_col(FieldArray::from_arr( + "price", + Array::from_decimal64(DecimalArray::::from_slice(&[10050], 18, 4)), + )); + + let mut narrower = Table::new_empty(); + narrower.add_col(FieldArray::from_arr( + "price", + Array::from_scalars(&[Scalar::Decimal64(20075, 12, 4)]), + )); + + assert!(typed.insert_rows(1, &narrower).is_err()); + } + #[cfg(feature = "chunked")] #[test] fn test_table_split_basic() { diff --git a/src/structs/variants/decimal.rs b/src/structs/variants/decimal.rs index 11ae184..5d757a0 100644 --- a/src/structs/variants/decimal.rs +++ b/src/structs/variants/decimal.rs @@ -57,6 +57,7 @@ use std::fmt::{Display, Formatter}; +use crate::enums::error::MinarrowError; use crate::enums::shape_dim::ShapeDim; use crate::traits::concatenate::Concatenate; use crate::traits::print::MAX_PREVIEW; @@ -749,16 +750,18 @@ impl Shape for DecimalArray { } // --------------------------------------------------------------------------- -// Concatenate - validates matching scale +// Concatenate - validates matching scale and precision // --------------------------------------------------------------------------- impl Concatenate for DecimalArray { - fn concat( - mut self, - other: Self, - ) -> core::result::Result { + /// Concatenates two decimal arrays of the same precision and scale. + /// + /// Precision is a column constraint, so a difference in either component + /// returns `IncompatibleTypeError` rather than widening to the larger + /// precision. + fn concat(mut self, other: Self) -> core::result::Result { if self.scale != other.scale { - return Err(crate::enums::error::MinarrowError::IncompatibleTypeError { + return Err(MinarrowError::IncompatibleTypeError { from: "DecimalArray", to: "DecimalArray", message: Some(format!( @@ -767,8 +770,16 @@ impl Concatenate for DecimalArray { )), }); } - // Take the wider precision to accommodate both operands - self.precision = self.precision.max(other.precision); + if self.precision != other.precision { + return Err(MinarrowError::IncompatibleTypeError { + from: "DecimalArray", + to: "DecimalArray", + message: Some(format!( + "precision mismatch: {} vs {}", + self.precision, other.precision + )), + }); + } self.append_array(&other); Ok(self) } @@ -1750,11 +1761,17 @@ mod tests { } #[test] - fn test_concat_takes_max_precision() { + fn test_concat_mismatched_precision_errors() { let arr1 = DecimalArray::::from_slice(&[100], 8, 2); let arr2 = DecimalArray::::from_slice(&[200], 12, 2); - let result = arr1.concat(arr2).unwrap(); - assert_eq!(result.precision, 12); + let err = arr1.clone().concat(arr2.clone()).unwrap_err(); + assert!( + format!("{}", err).contains("precision mismatch"), + "Expected precision mismatch error, got: {}", + err + ); + // The mismatch is symmetric: the wider precision first is also an error. + assert!(arr2.concat(arr1).is_err()); } #[test] diff --git a/src/traits/byte_size.rs b/src/traits/byte_size.rs index 123ebfa..0c45f00 100644 --- a/src/traits/byte_size.rs +++ b/src/traits/byte_size.rs @@ -923,11 +923,11 @@ impl ByteSize for Scalar { #[cfg(feature = "datetime")] Scalar::Interval => 0, #[cfg(feature = "decimal")] - Scalar::Decimal32(_, _) => size_of::(), + Scalar::Decimal32(_, _, _) => size_of::(), #[cfg(feature = "decimal")] - Scalar::Decimal64(_, _) => size_of::(), + Scalar::Decimal64(_, _, _) => size_of::(), #[cfg(feature = "decimal")] - Scalar::Decimal128(_, _) => size_of::(), + Scalar::Decimal128(_, _, _) => size_of::(), } } }