Skip to main content

competitive/num/decimal/
addsub.rs

1use super::*;
2use std::{
3    mem::replace,
4    ops::{Add, AddAssign, Sub, SubAssign},
5};
6
7fn add_carry(carry: bool, lhs: u64, rhs: u64, out: &mut u64) -> bool {
8    let mut sum = lhs + rhs + carry as u64;
9    let cond = sum >= RADIX;
10    if cond {
11        sum -= RADIX;
12    }
13    *out = sum;
14    cond
15}
16
17fn add_absolute_parts(lhs: &mut Decimal, rhs: &Decimal) {
18    let mut carry = false;
19
20    // decimal part
21    let lhs_decimal_len = lhs.decimal.len();
22    if lhs_decimal_len < rhs.decimal.len() {
23        for (l, r) in lhs
24            .decimal
25            .iter_mut()
26            .rev()
27            .zip(rhs.decimal[..lhs_decimal_len].iter().rev())
28        {
29            carry = add_carry(carry, *l, *r, l);
30        }
31        lhs.decimal
32            .extend_from_slice(&rhs.decimal[lhs_decimal_len..]);
33    } else {
34        for (l, r) in lhs.decimal[..rhs.decimal.len()]
35            .iter_mut()
36            .rev()
37            .zip(rhs.decimal.iter().rev())
38        {
39            carry = add_carry(carry, *l, *r, l);
40        }
41    }
42
43    // integer part
44    let lhs_integer_len = lhs.integer.len();
45    if lhs_integer_len < rhs.integer.len() {
46        for (l, r) in lhs.integer.iter_mut().zip(&rhs.integer[..lhs_integer_len]) {
47            carry = add_carry(carry, *l, *r, l);
48        }
49        lhs.integer
50            .extend_from_slice(&rhs.integer[lhs_integer_len..]);
51        if carry {
52            for l in lhs.integer[lhs_integer_len..].iter_mut() {
53                carry = add_carry(carry, *l, 0, l);
54                if !carry {
55                    break;
56                }
57            }
58        }
59    } else {
60        for (l, r) in lhs.integer.iter_mut().zip(&rhs.integer) {
61            carry = add_carry(carry, *l, *r, l);
62        }
63        if carry {
64            for l in lhs.integer[rhs.integer.len()..].iter_mut() {
65                carry = add_carry(carry, *l, 0, l);
66                if !carry {
67                    break;
68                }
69            }
70        }
71    }
72
73    if carry {
74        lhs.integer.push(carry as u64);
75    }
76
77    lhs.normalize();
78}
79
80fn sub_borrow(borrow: bool, lhs: u64, rhs: u64, out: &mut u64) -> bool {
81    let (sum, borrow1) = lhs.overflowing_sub(rhs);
82    let (mut sum, borrow2) = sum.overflowing_sub(borrow as u64);
83    let borrow = borrow1 || borrow2;
84    if borrow {
85        sum = sum.wrapping_add(RADIX);
86    }
87    *out = sum;
88    borrow
89}
90
91// assume |lhs| >= |rhs|
92fn sub_absolute_parts_gte(lhs: &Decimal, rhs: &mut Decimal) {
93    debug_assert!(matches!(lhs.cmp_absolute_parts(rhs), Ordering::Greater));
94
95    let mut borrow = false;
96
97    // decimal part
98    let rhs_decimal_len = rhs.decimal.len();
99    if lhs.decimal.len() > rhs_decimal_len {
100        for (l, r) in lhs.decimal[..rhs_decimal_len]
101            .iter()
102            .rev()
103            .zip(rhs.decimal.iter_mut().rev())
104        {
105            borrow = sub_borrow(borrow, *l, *r, r);
106        }
107        rhs.decimal
108            .extend_from_slice(&lhs.decimal[rhs_decimal_len..]);
109    } else {
110        for r in rhs.decimal[lhs.decimal.len()..].iter_mut().rev() {
111            borrow = sub_borrow(borrow, 0, *r, r);
112        }
113        for (l, r) in lhs
114            .decimal
115            .iter()
116            .rev()
117            .zip(rhs.decimal[..lhs.decimal.len()].iter_mut().rev())
118        {
119            borrow = sub_borrow(borrow, *l, *r, r);
120        }
121    }
122
123    // integer part
124    let rhs_integer_len = rhs.integer.len();
125    if lhs.integer.len() > rhs_integer_len {
126        for (l, r) in lhs.integer[..rhs_integer_len]
127            .iter()
128            .zip(rhs.integer.iter_mut())
129        {
130            borrow = sub_borrow(borrow, *l, *r, r);
131        }
132        rhs.integer
133            .extend_from_slice(&lhs.integer[rhs_integer_len..]);
134        if borrow {
135            for r in rhs.integer[rhs_integer_len..].iter_mut() {
136                borrow = sub_borrow(borrow, *r, 0, r);
137                if !borrow {
138                    break;
139                }
140            }
141        }
142    } else {
143        debug_assert_eq!(lhs.integer.len(), rhs_integer_len);
144        for (l, r) in lhs.integer.iter().zip(&mut rhs.integer) {
145            borrow = sub_borrow(borrow, *l, *r, r);
146        }
147    }
148
149    assert!(
150        !borrow,
151        "Cannot subtract lhs from rhs because lhs is smaller than rhs"
152    );
153
154    rhs.normalize();
155}
156
157macro_rules! add {
158    ($lhs:expr, $lhs_owned:expr, $rhs:expr, $rhs_owned:expr) => {
159        match ($lhs.sign, $rhs.sign) {
160            (Sign::Zero, _) => $rhs_owned,
161            (_, Sign::Zero) => $lhs_owned,
162            (Sign::Plus, Sign::Plus) | (Sign::Minus, Sign::Minus) => {
163                let mut lhs = $lhs_owned;
164                add_absolute_parts(&mut lhs, &$rhs);
165                lhs
166            }
167            (Sign::Plus, Sign::Minus) | (Sign::Minus, Sign::Plus) => {
168                match $lhs.cmp_absolute_parts(&$rhs) {
169                    Ordering::Less => {
170                        let mut lhs = $lhs_owned;
171                        sub_absolute_parts_gte(&$rhs, &mut lhs);
172                        lhs.sign = $rhs.sign;
173                        lhs
174                    }
175                    Ordering::Equal => ZERO,
176                    Ordering::Greater => {
177                        let mut rhs = $rhs_owned;
178                        sub_absolute_parts_gte(&$lhs, &mut rhs);
179                        rhs.sign = $lhs.sign;
180                        rhs
181                    }
182                }
183            }
184        }
185    };
186}
187
188macro_rules! sub {
189    ($lhs:expr, $lhs_owned:expr, $rhs:expr, $rhs_owned:expr) => {
190        match ($lhs.sign, $rhs.sign) {
191            (Sign::Zero, _) => -$rhs_owned,
192            (_, Sign::Zero) => $lhs_owned,
193            (Sign::Plus, Sign::Minus) | (Sign::Minus, Sign::Plus) => {
194                let mut lhs = $lhs_owned;
195                add_absolute_parts(&mut lhs, &$rhs);
196                lhs
197            }
198            (Sign::Plus, Sign::Plus) | (Sign::Minus, Sign::Minus) => {
199                match $lhs.cmp_absolute_parts(&$rhs) {
200                    Ordering::Less => {
201                        let mut lhs = $lhs_owned;
202                        sub_absolute_parts_gte(&$rhs, &mut lhs);
203                        lhs.sign = -$rhs.sign;
204                        lhs
205                    }
206                    Ordering::Equal => ZERO,
207                    Ordering::Greater => {
208                        let mut rhs = $rhs_owned;
209                        sub_absolute_parts_gte(&$lhs, &mut rhs);
210                        rhs
211                    }
212                }
213            }
214        }
215    };
216}
217
218macro_rules! impl_binop {
219    (impl $Trait:ident for Decimal, $method:ident, $macro:ident) => {
220        impl $Trait<Decimal> for Decimal {
221            type Output = Decimal;
222
223            fn $method(self, rhs: Decimal) -> Self::Output {
224                $macro!(self, self, rhs, rhs)
225            }
226        }
227
228        impl $Trait<&Decimal> for Decimal {
229            type Output = Decimal;
230
231            fn $method(self, rhs: &Decimal) -> Self::Output {
232                $macro!(self, self, rhs, rhs.clone())
233            }
234        }
235
236        impl $Trait<Decimal> for &Decimal {
237            type Output = Decimal;
238
239            fn $method(self, rhs: Decimal) -> Self::Output {
240                $macro!(self, self.clone(), rhs, rhs)
241            }
242        }
243
244        impl $Trait<&Decimal> for &Decimal {
245            type Output = Decimal;
246
247            fn $method(self, rhs: &Decimal) -> Self::Output {
248                $macro!(self, self.clone(), rhs, rhs.clone())
249            }
250        }
251    };
252}
253impl_binop!(impl Add for Decimal, add, add);
254impl_binop!(impl Sub for Decimal, sub, sub);
255
256macro_rules! impl_binop_assign {
257    (impl $Trait:ident for Decimal, $method:ident, $op:tt) => {
258        impl $Trait for Decimal {
259            fn $method(&mut self, rhs: Decimal) {
260                let lhs = replace(self, ZERO);
261                *self = lhs $op rhs;
262            }
263        }
264
265        impl $Trait<&Decimal> for Decimal {
266            fn $method(&mut self, rhs: &Decimal) {
267                let lhs = replace(self, ZERO);
268                *self = lhs $op rhs;
269            }
270        }
271    };
272}
273
274impl_binop_assign!(impl AddAssign for Decimal, add_assign, +);
275impl_binop_assign!(impl SubAssign for Decimal, sub_assign, -);