1use std::cell::Cell;
2
3use polars_arrow::bitmap::Bitmap;
4use polars_arrow::compute::utils::combine_validities_and;
5use polars_compute::decimal::{
6 DEC128_MAX_PREC, dec128_add_scaled, dec128_div_scaled, dec128_int_div_scaled,
7 dec128_mul_scaled, dec128_rem_scaled, dec128_rescale, dec128_sub_scaled, i64_binary_values,
8 i64_unary_values,
9};
10
11use super::*;
12use crate::prelude::arity::{
13 apply_binary_kernel_broadcast, broadcast_binary_elementwise_values,
14 broadcast_try_binary_elementwise,
15};
16
17impl DecimalChunked {
18 fn apply_scaled_kernel(
22 &self,
23 rhs: &Self,
24 scale: usize,
25 op: &str,
26 kernel: impl Fn(i128, usize, i128, usize, usize) -> Option<i128>,
27 ) -> PolarsResult<Self> {
28 let left_s = self.scale();
29 let right_s = rhs.scale();
30 let phys = broadcast_try_binary_elementwise(
31 self.physical(),
32 rhs.physical(),
33 |opt_l, opt_r| {
34 let (Some(l), Some(r)) = (opt_l, opt_r) else {
35 return PolarsResult::Ok(None);
36 };
37 let ret = kernel(l, left_s, r, right_s, scale).ok_or_else(|| {
38 if r == 0 {
40 polars_err!(ComputeError: "division by zero Decimal")
41 } else {
42 polars_err!(
43 ComputeError: "overflow in decimal {op}: result doesn't fit Decimal({DEC128_MAX_PREC}, {scale})"
44 )
45 }
46 })?;
47 Ok(Some(ret))
48 },
49 )?;
50 Ok(phys.into_decimal_unchecked(DEC128_MAX_PREC, scale))
51 }
52
53 fn apply_scaled_kernel_values(
57 &self,
58 rhs: &Self,
59 scale: usize,
60 op: &str,
61 kernel: impl Fn(i128, usize, i128, usize, usize) -> Option<i128>,
62 ) -> PolarsResult<Self> {
63 let left_s = self.scale();
64 let right_s = rhs.scale();
65 let mut failed = false;
66 let phys = broadcast_binary_elementwise_values(self.physical(), rhs.physical(), |l, r| {
67 let ret = kernel(l, left_s, r, right_s, scale);
68 failed |= ret.is_none();
69 ret.unwrap_or(0)
70 });
71 if failed {
72 return self.apply_scaled_kernel(rhs, scale, op, kernel);
73 }
74 Ok(phys.into_decimal_unchecked(DEC128_MAX_PREC, scale))
75 }
76
77 fn apply_i64_values(
80 &self,
81 rhs: &Self,
82 scale: usize,
83 op: impl Fn(i64, i64) -> i128,
84 ) -> Option<Self> {
85 let failed = Cell::new(false);
86 let to_arr = |values: Option<Vec<i128>>, validity: Option<Bitmap>| {
87 let Some(values) = values else {
88 failed.set(true);
90 return Int128Array::new_empty(ArrowDataType::Int128);
91 };
92 Int128Array::from_vec(values).with_validity(validity)
93 };
94 let phys = apply_binary_kernel_broadcast(
95 self.physical(),
96 rhs.physical(),
97 |l, r| {
98 let values = i64_binary_values(l.values(), r.values(), &op);
99 to_arr(values, combine_validities_and(l.validity(), r.validity()))
100 },
101 |l, r| {
102 let values = i64::try_from(l)
103 .ok()
104 .and_then(|l| i64_unary_values(r.values(), |r| op(l, r)));
105 to_arr(values, r.validity().cloned())
106 },
107 |l, r| {
108 let values = i64::try_from(r)
109 .ok()
110 .and_then(|r| i64_unary_values(l.values(), |l| op(l, r)));
111 to_arr(values, l.validity().cloned())
112 },
113 );
114 (!failed.get()).then(|| phys.into_decimal_unchecked(DEC128_MAX_PREC, scale))
115 }
116
117 fn scalar_with_scale(&self, scale: usize) -> Option<Self> {
119 if self.len() != 1 || self.scale() == scale {
120 return None;
121 }
122 let value = dec128_rescale(
123 self.physical().get(0)?,
124 self.scale(),
125 DEC128_MAX_PREC,
126 scale,
127 )?;
128 Some(
129 Int128Chunked::from_slice(self.name().clone(), &[value])
130 .into_decimal_unchecked(DEC128_MAX_PREC, scale),
131 )
132 }
133
134 fn add_sub(
137 &self,
138 rhs: &Self,
139 op: &str,
140 i64_op: impl Fn(i64, i64) -> i128,
141 kernel: impl Fn(i128, usize, i128, usize, usize) -> Option<i128>,
142 ) -> PolarsResult<Self> {
143 let scale = self.scale().max(rhs.scale());
144 let lhs_scalar = self.scalar_with_scale(scale);
145 let rhs_scalar = rhs.scalar_with_scale(scale);
146 let lhs = lhs_scalar.as_ref().unwrap_or(self);
147 let rhs = rhs_scalar.as_ref().unwrap_or(rhs);
148 if lhs.scale() == rhs.scale()
149 && let Some(out) = lhs.apply_i64_values(rhs, scale, i64_op)
150 {
151 return Ok(out);
152 }
153 lhs.apply_scaled_kernel_values(rhs, scale, op, kernel)
154 }
155
156 pub fn mul_with_scale(&self, rhs: &Self, scale: usize) -> PolarsResult<Self> {
158 if self.scale() + rhs.scale() == scale
159 && let Some(out) = self.apply_i64_values(rhs, scale, |l, r| l as i128 * r as i128)
160 {
161 return Ok(out);
162 }
163 self.apply_scaled_kernel_values(rhs, scale, "multiplication", dec128_mul_scaled)
164 }
165
166 pub fn div_with_scale(&self, rhs: &Self, scale: usize) -> PolarsResult<Self> {
168 self.apply_scaled_kernel(rhs, scale, "division", dec128_div_scaled)
169 }
170
171 pub fn rem_with(&self, rhs: &Self, floor: bool) -> PolarsResult<Self> {
173 let scale = self.scale().max(rhs.scale());
174 self.apply_scaled_kernel(rhs, scale, "remainder", |l, sl, r, sr, s| {
175 dec128_rem_scaled(l, sl, r, sr, s, floor)
176 })
177 }
178
179 pub fn int_div(&self, rhs: &Self, floor: bool) -> PolarsResult<Self> {
181 self.int_div_with_scale(rhs, self.scale().max(rhs.scale()), floor)
182 }
183
184 pub fn int_div_with_scale(&self, rhs: &Self, scale: usize, floor: bool) -> PolarsResult<Self> {
186 self.apply_scaled_kernel(rhs, scale, "integer division", |l, sl, r, sr, s| {
187 dec128_int_div_scaled(l, sl, r, sr, s, floor)
188 })
189 }
190}
191
192impl Add for &DecimalChunked {
193 type Output = PolarsResult<DecimalChunked>;
194
195 fn add(self, rhs: Self) -> Self::Output {
196 self.add_sub(
197 rhs,
198 "addition",
199 |l, r| l as i128 + r as i128,
200 dec128_add_scaled,
201 )
202 }
203}
204
205impl Sub for &DecimalChunked {
206 type Output = PolarsResult<DecimalChunked>;
207
208 fn sub(self, rhs: Self) -> Self::Output {
209 self.add_sub(
210 rhs,
211 "subtraction",
212 |l, r| l as i128 - r as i128,
213 dec128_sub_scaled,
214 )
215 }
216}
217
218impl Mul for &DecimalChunked {
219 type Output = PolarsResult<DecimalChunked>;
220
221 fn mul(self, rhs: Self) -> Self::Output {
222 self.mul_with_scale(rhs, self.scale().max(rhs.scale()))
223 }
224}
225
226impl Div for &DecimalChunked {
227 type Output = PolarsResult<DecimalChunked>;
228
229 fn div(self, rhs: Self) -> Self::Output {
230 self.div_with_scale(rhs, self.scale().max(rhs.scale()))
231 }
232}