Skip to main content

pyo3/conversions/
bigdecimal.rs

1#![cfg(feature = "bigdecimal")]
2//! Conversions to and from [bigdecimal](https://docs.rs/bigdecimal)'s [`BigDecimal`] type.
3//!
4//! This is useful for converting Python's decimal.Decimal into and from a native Rust type.
5//!
6//! # Setup
7//!
8//! To use this feature, add to your **`Cargo.toml`**:
9//!
10//! ```toml
11//! [dependencies]
12#![doc = concat!("pyo3 = { version = \"", env!("CARGO_PKG_VERSION"),  "\", features = [\"bigdecimal\"] }")]
13//! bigdecimal = "0.4"
14//! ```
15//!
16//! Note that you must use a compatible version of bigdecimal and PyO3.
17//! The required bigdecimal version may vary based on the version of PyO3.
18//!
19//! # Example
20//!
21//! Rust code to create a function that adds one to a BigDecimal
22//!
23//! ```rust
24//! use bigdecimal::BigDecimal;
25//! use pyo3::prelude::*;
26//!
27//! #[pyfunction]
28//! fn add_one(d: BigDecimal) -> BigDecimal {
29//!     d + 1
30//! }
31//!
32//! #[pymodule]
33//! # #[pyo3(name = "example_bigdecimal")]
34//! fn my_module(m: &Bound<'_, PyModule>) -> PyResult<()> {
35//!     m.add_function(wrap_pyfunction!(add_one, m)?)?;
36//!     Ok(())
37//! }
38//! ```
39//!
40//! Python code that validates the functionality
41//!
42//!
43//! ```python
44//! from my_module import add_one
45//! from decimal import Decimal
46//!
47//! d = Decimal("2")
48//! value = add_one(d)
49//!
50//! assert d + 1 == value
51//! ```
52
53use core::str::FromStr;
54
55#[cfg(feature = "experimental-inspect")]
56use crate::inspect::PyStaticExpr;
57use crate::platform::prelude::*;
58#[cfg(feature = "experimental-inspect")]
59use crate::type_hint_identifier;
60use crate::types::PyTuple;
61use crate::{
62    Borrowed, Bound, FromPyObject, IntoPyObject, Py, PyAny, PyErr, PyResult, Python,
63    exceptions::PyValueError,
64    sync::PyOnceLock,
65    types::{PyAnyMethods, PyStringMethods, PyType},
66};
67use bigdecimal::BigDecimal;
68use num_bigint::Sign;
69
70fn get_decimal_cls(py: Python<'_>) -> PyResult<&Bound<'_, PyType>> {
71    static DECIMAL_CLS: PyOnceLock<Py<PyType>> = PyOnceLock::new();
72    DECIMAL_CLS.import(py, "decimal", "Decimal")
73}
74
75fn get_invalid_operation_error_cls(py: Python<'_>) -> PyResult<&Bound<'_, PyType>> {
76    static INVALID_OPERATION_CLS: PyOnceLock<Py<PyType>> = PyOnceLock::new();
77    INVALID_OPERATION_CLS.import(py, "decimal", "InvalidOperation")
78}
79
80impl FromPyObject<'_, '_> for BigDecimal {
81    type Error = PyErr;
82
83    #[cfg(feature = "experimental-inspect")]
84    const INPUT_TYPE: PyStaticExpr = type_hint_identifier!("decimal", "Decimal");
85
86    fn extract(obj: Borrowed<'_, '_, PyAny>) -> PyResult<Self> {
87        let py_str = &obj.str()?;
88        let rs_str = &py_str.to_cow()?;
89        BigDecimal::from_str(rs_str).map_err(|e| PyValueError::new_err(e.to_string()))
90    }
91}
92
93impl<'py> IntoPyObject<'py> for BigDecimal {
94    type Target = PyAny;
95
96    type Output = Bound<'py, Self::Target>;
97
98    type Error = PyErr;
99
100    #[cfg(feature = "experimental-inspect")]
101    const OUTPUT_TYPE: PyStaticExpr = type_hint_identifier!("decimal", "Decimal");
102
103    fn into_pyobject(self, py: Python<'py>) -> Result<Self::Output, Self::Error> {
104        let cls = get_decimal_cls(py)?;
105        let (bigint, scale) = self.into_bigint_and_scale();
106        if scale == 0 {
107            return cls.call1((bigint,));
108        }
109        let exponent = scale.checked_neg().ok_or_else(|| {
110            get_invalid_operation_error_cls(py)
111                .map_or_else(|err| err, |cls| PyErr::from_type(cls.clone(), ()))
112        })?;
113        let (sign, digits) = bigint.to_radix_be(10);
114        let signed = matches!(sign, Sign::Minus).into_pyobject(py)?;
115        let digits = PyTuple::new(py, digits)?;
116
117        cls.call1(((signed, digits, exponent),))
118    }
119}
120
121#[cfg(test)]
122mod test_bigdecimal {
123    use super::*;
124    use crate::types::PyDict;
125    use crate::types::dict::PyDictMethods;
126    use alloc::ffi::CString;
127
128    use bigdecimal::{One, Zero};
129    #[cfg(not(target_arch = "wasm32"))]
130    use proptest::prelude::*;
131
132    macro_rules! convert_constants {
133        ($name:ident, $rs:expr, $py:literal) => {
134            #[test]
135            fn $name() {
136                Python::attach(|py| {
137                    let rs_orig = $rs;
138                    let rs_dec = rs_orig.clone().into_pyobject(py).unwrap();
139                    let locals = PyDict::new(py);
140                    locals.set_item("rs_dec", &rs_dec).unwrap();
141                    // Checks if BigDecimal -> Python Decimal conversion is correct
142                    py.run(
143                        &CString::new(format!(
144                            "import decimal\npy_dec = decimal.Decimal(\"{}\")\nassert py_dec == rs_dec",
145                            $py
146                        ))
147                        .unwrap(),
148                        None,
149                        Some(&locals),
150                    )
151                    .unwrap();
152                    // Checks if Python Decimal -> BigDecimal conversion is correct
153                    let py_dec = locals.get_item("py_dec").unwrap().unwrap();
154                    let py_result: BigDecimal = py_dec.extract().unwrap();
155                    assert_eq!(rs_orig, py_result);
156                })
157            }
158        };
159    }
160
161    convert_constants!(convert_zero, BigDecimal::zero(), "0");
162    convert_constants!(convert_one, BigDecimal::one(), "1");
163    convert_constants!(convert_neg_one, -BigDecimal::one(), "-1");
164    convert_constants!(convert_two, BigDecimal::from(2), "2");
165    convert_constants!(convert_ten, BigDecimal::from_str("10").unwrap(), "10");
166    convert_constants!(
167        convert_one_hundred_point_one,
168        BigDecimal::from_str("100.1").unwrap(),
169        "100.1"
170    );
171    convert_constants!(
172        convert_one_thousand,
173        BigDecimal::from_str("1000").unwrap(),
174        "1000"
175    );
176    convert_constants!(
177        convert_scientific,
178        BigDecimal::from_str("1e10").unwrap(),
179        "1e10"
180    );
181
182    #[cfg(not(target_arch = "wasm32"))]
183    proptest! {
184        #[test]
185        fn test_roundtrip(
186            number in 0..28u32
187        ) {
188            let num = BigDecimal::from(number);
189            Python::attach(|py| {
190                let rs_dec = num.clone().into_pyobject(py).unwrap();
191                let locals = PyDict::new(py);
192                locals.set_item("rs_dec", &rs_dec).unwrap();
193                py.run(
194                    &CString::new(format!(
195                       "import decimal\npy_dec = decimal.Decimal(\"{num}\")\nassert py_dec == rs_dec")).unwrap(),
196                None, Some(&locals)).unwrap();
197                let roundtripped: BigDecimal = rs_dec.extract().unwrap();
198                assert_eq!(num, roundtripped);
199            })
200        }
201
202        #[test]
203        fn test_integers(num in any::<i64>()) {
204            Python::attach(|py| {
205                let py_num = num.into_pyobject(py).unwrap();
206                let roundtripped: BigDecimal = py_num.extract().unwrap();
207                let rs_dec = BigDecimal::from(num);
208                assert_eq!(rs_dec, roundtripped);
209            })
210        }
211    }
212
213    #[test]
214    fn test_nan() {
215        Python::attach(|py| {
216            let locals = PyDict::new(py);
217            py.run(
218                c"import decimal\npy_dec = decimal.Decimal(\"NaN\")",
219                None,
220                Some(&locals),
221            )
222            .unwrap();
223            let py_dec = locals.get_item("py_dec").unwrap().unwrap();
224            let roundtripped: Result<BigDecimal, PyErr> = py_dec.extract();
225            assert!(roundtripped.is_err());
226        })
227    }
228
229    #[test]
230    fn test_infinity() {
231        Python::attach(|py| {
232            let locals = PyDict::new(py);
233            py.run(
234                c"import decimal\npy_dec = decimal.Decimal(\"Infinity\")",
235                None,
236                Some(&locals),
237            )
238            .unwrap();
239            let py_dec = locals.get_item("py_dec").unwrap().unwrap();
240            let roundtripped: Result<BigDecimal, PyErr> = py_dec.extract();
241            assert!(roundtripped.is_err());
242        })
243    }
244
245    #[test]
246    fn test_no_precision_loss() {
247        Python::attach(|py| {
248            let src = "1e4";
249            let expected = get_decimal_cls(py)
250                .unwrap()
251                .call1((src,))
252                .unwrap()
253                .call_method0("as_tuple")
254                .unwrap();
255            let actual = src
256                .parse::<BigDecimal>()
257                .unwrap()
258                .into_pyobject(py)
259                .unwrap()
260                .call_method0("as_tuple")
261                .unwrap();
262
263            assert!(actual.eq(expected).unwrap());
264        });
265    }
266}
⚠️ Internal Docs ⚠️ Not Public API 👉 Official Docs Here