pyo3/conversions/
bigdecimal.rs1#![cfg(feature = "bigdecimal")]
2#![doc = concat!("pyo3 = { version = \"", env!("CARGO_PKG_VERSION"), "\", features = [\"bigdecimal\"] }")]
13use 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 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 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}