33
44//! Native execution of the arithmetic operators over decimal arrays.
55//!
6- //! Both operands share a logical [`DecimalDType`] (equal precision and scale) and the result
7- //! keeps that dtype: fixed-point arithmetic at the shared scale. Add and Sub apply directly to
8- //! the unscaled stored integers and are exact; Mul and Div require rescaling and are not yet
9- //! implemented.
6+ //! Both operands share a logical [`DecimalDType`] (equal precision and scale). Add and Sub apply
7+ //! directly to the unscaled stored integers and are exact at that shared scale. The result reserves
8+ //! one additional precision digit for a carry, capped at Vortex's maximum decimal precision. Mul
9+ //! and Div require rescaling and are not yet implemented.
1010//!
1111//! Lanes execute in a working width chosen so that in-precision inputs cannot spuriously
1212//! overflow an intermediate value. An operation that overflows the result precision on a valid
@@ -25,6 +25,7 @@ use super::CheckedValues;
2525use super :: check_numeric_errors;
2626use super :: checked_all_lanes;
2727use super :: checked_valid_lanes;
28+ use super :: decimal_add_sub_result_dtype;
2829use crate :: ArrayRef ;
2930use crate :: ExecutionCtx ;
3031use crate :: IntoArray ;
@@ -54,10 +55,11 @@ pub(super) fn execute_numeric_decimal(
5455 let DType :: Decimal ( decimal_dtype, _) = lhs. dtype ( ) else {
5556 vortex_bail ! ( "expected a decimal dtype, got {}" , lhs. dtype( ) ) ;
5657 } ;
57- let decimal_dtype = * decimal_dtype;
58- let result_dtype = lhs
59- . dtype ( )
60- . with_nullability ( lhs. dtype ( ) . nullability ( ) | rhs. dtype ( ) . nullability ( ) ) ;
58+ let result_decimal_dtype = decimal_add_sub_result_dtype ( * decimal_dtype) ;
59+ let result_dtype = DType :: Decimal (
60+ result_decimal_dtype,
61+ lhs. dtype ( ) . nullability ( ) | rhs. dtype ( ) . nullability ( ) ,
62+ ) ;
6163
6264 let lhs = DecimalOperand :: try_new ( lhs, ctx) ?;
6365 let rhs = DecimalOperand :: try_new ( rhs, ctx) ?;
@@ -67,82 +69,61 @@ pub(super) fn execute_numeric_decimal(
6769 let validity = lhs. validity ( ) . and ( rhs. validity ( ) ) ?;
6870 let valid_rows = validity. execute_mask ( len, ctx) ?;
6971
70- let work = working_type ( decimal_dtype) ;
71- let output = DecimalType :: smallest_decimal_value_type ( & decimal_dtype) ;
72- match ( work, output) {
73- ( DecimalType :: I8 , DecimalType :: I8 ) => execute_decimal_at_widths :: < i8 , i8 > (
72+ match DecimalType :: smallest_decimal_value_type ( & result_decimal_dtype) {
73+ DecimalType :: I8 => execute_decimal_at_widths :: < i8 , i8 > (
7474 & lhs,
7575 & rhs,
7676 op,
77- decimal_dtype ,
77+ result_decimal_dtype ,
7878 & result_dtype,
7979 validity,
8080 & valid_rows,
8181 ) ,
82- ( DecimalType :: I16 , DecimalType :: I8 ) => execute_decimal_at_widths :: < i16 , i8 > (
82+ DecimalType :: I16 => execute_decimal_at_widths :: < i16 , i16 > (
8383 & lhs,
8484 & rhs,
8585 op,
86- decimal_dtype ,
86+ result_decimal_dtype ,
8787 & result_dtype,
8888 validity,
8989 & valid_rows,
9090 ) ,
91- ( DecimalType :: I16 , DecimalType :: I16 ) => execute_decimal_at_widths :: < i16 , i16 > (
91+ DecimalType :: I32 => execute_decimal_at_widths :: < i32 , i32 > (
9292 & lhs,
9393 & rhs,
9494 op,
95- decimal_dtype ,
95+ result_decimal_dtype ,
9696 & result_dtype,
9797 validity,
9898 & valid_rows,
9999 ) ,
100- ( DecimalType :: I32 , DecimalType :: I32 ) => execute_decimal_at_widths :: < i32 , i32 > (
100+ DecimalType :: I64 => execute_decimal_at_widths :: < i64 , i64 > (
101101 & lhs,
102102 & rhs,
103103 op,
104- decimal_dtype ,
104+ result_decimal_dtype ,
105105 & result_dtype,
106106 validity,
107107 & valid_rows,
108108 ) ,
109- ( DecimalType :: I64 , DecimalType :: I64 ) => execute_decimal_at_widths :: < i64 , i64 > (
109+ DecimalType :: I128 => execute_decimal_at_widths :: < i128 , i128 > (
110110 & lhs,
111111 & rhs,
112112 op,
113- decimal_dtype ,
113+ result_decimal_dtype ,
114114 & result_dtype,
115115 validity,
116116 & valid_rows,
117117 ) ,
118- ( DecimalType :: I128 , DecimalType :: I128 ) => execute_decimal_at_widths :: < i128 , i128 > (
118+ DecimalType :: I256 => execute_decimal_at_widths :: < i256 , i256 > (
119119 & lhs,
120120 & rhs,
121121 op,
122- decimal_dtype ,
122+ result_decimal_dtype ,
123123 & result_dtype,
124124 validity,
125125 & valid_rows,
126126 ) ,
127- ( DecimalType :: I256 , DecimalType :: I128 ) => execute_decimal_at_widths :: < i256 , i128 > (
128- & lhs,
129- & rhs,
130- op,
131- decimal_dtype,
132- & result_dtype,
133- validity,
134- & valid_rows,
135- ) ,
136- ( DecimalType :: I256 , DecimalType :: I256 ) => execute_decimal_at_widths :: < i256 , i256 > (
137- & lhs,
138- & rhs,
139- op,
140- decimal_dtype,
141- & result_dtype,
142- validity,
143- & valid_rows,
144- ) ,
145- _ => vortex_bail ! ( "unsupported decimal working/output width combination: {work}/{output}" ) ,
146127 }
147128}
148129
@@ -198,33 +179,6 @@ impl DecimalOperand {
198179 }
199180}
200181
201- /// Choose the smallest lane width that can represent every sum or difference of two valid inputs.
202- fn working_type ( dtype : DecimalDType ) -> DecimalType {
203- let precision = dtype. precision ( ) as usize ;
204- let max = <i256 as NativeDecimalType >:: MAX_BY_PRECISION [ precision] ;
205- let max_result = max
206- . checked_add ( & max)
207- . vortex_expect ( "the sum of two valid decimal values must fit in i256" ) ;
208- smallest_value_type ( & DecimalValue :: from ( max_result) )
209- }
210-
211- /// The smallest decimal value type that can represent `value`, regardless of its stored width.
212- fn smallest_value_type ( value : & DecimalValue ) -> DecimalType {
213- if value. cast :: < i8 > ( ) . is_some ( ) {
214- DecimalType :: I8
215- } else if value. cast :: < i16 > ( ) . is_some ( ) {
216- DecimalType :: I16
217- } else if value. cast :: < i32 > ( ) . is_some ( ) {
218- DecimalType :: I32
219- } else if value. cast :: < i64 > ( ) . is_some ( ) {
220- DecimalType :: I64
221- } else if value. cast :: < i128 > ( ) . is_some ( ) {
222- DecimalType :: I128
223- } else {
224- DecimalType :: I256
225- }
226- }
227-
228182/// Per-execution constants for checked decimal lane operations at working width `W`.
229183struct DecimalOpPlan < W > {
230184 /// Inclusive stored-value bounds implied by the result precision.
@@ -289,7 +243,7 @@ fn execute_decimal_at_widths<W, O>(
289243 lhs : & DecimalOperand ,
290244 rhs : & DecimalOperand ,
291245 op : NumericOperator ,
292- decimal_dtype : DecimalDType ,
246+ result_decimal_dtype : DecimalDType ,
293247 result_dtype : & DType ,
294248 validity : Validity ,
295249 valid_rows : & Mask ,
@@ -303,15 +257,15 @@ where
303257 NumericOperator :: Add => execute_decimal_typed :: < W , O , DecimalAdd > (
304258 lhs,
305259 rhs,
306- decimal_dtype ,
260+ result_decimal_dtype ,
307261 result_dtype,
308262 validity,
309263 valid_rows,
310264 ) ,
311265 NumericOperator :: Sub => execute_decimal_typed :: < W , O , DecimalSub > (
312266 lhs,
313267 rhs,
314- decimal_dtype ,
268+ result_decimal_dtype ,
315269 result_dtype,
316270 validity,
317271 valid_rows,
@@ -326,7 +280,7 @@ where
326280fn execute_decimal_typed < W , O , Op > (
327281 lhs : & DecimalOperand ,
328282 rhs : & DecimalOperand ,
329- decimal_dtype : DecimalDType ,
283+ result_decimal_dtype : DecimalDType ,
330284 result_dtype : & DType ,
331285 validity : Validity ,
332286 valid_rows : & Mask ,
@@ -338,7 +292,7 @@ where
338292 Op : CheckedDecimalOp ,
339293{
340294 let len = lhs. len ( ) ;
341- let plan = DecimalOpPlan :: < W > :: new ( decimal_dtype ) ;
295+ let plan = DecimalOpPlan :: < W > :: new ( result_decimal_dtype ) ;
342296
343297 let checked = match ( lhs, rhs) {
344298 ( DecimalOperand :: Array { values : lhs, .. } , DecimalOperand :: Array { values : rhs, .. } ) => {
@@ -374,7 +328,7 @@ where
374328 return Ok ( ConstantArray :: new (
375329 Scalar :: decimal (
376330 DecimalValue :: from ( value) ,
377- decimal_dtype ,
331+ result_decimal_dtype ,
378332 result_dtype. nullability ( ) ,
379333 ) ,
380334 len,
@@ -389,7 +343,7 @@ where
389343
390344 Ok ( DecimalArray :: new (
391345 checked. values ,
392- decimal_dtype ,
346+ result_decimal_dtype ,
393347 validity. union_nullability ( result_dtype. nullability ( ) ) ,
394348 )
395349 . into_array ( ) )
0 commit comments