1use std::cmp::Ordering;
21use std::fmt;
22
23use serde::{Deserialize, Deserializer, Serialize, Serializer};
24
25#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
27pub struct Decimal {
28 unscaled: i128,
29 scale: u32,
30}
31
32const MAX_PRECISION: u32 = 38;
34
35impl Decimal {
36 pub fn from_parts(unscaled: i128, scale: u32) -> Option<Self> {
39 (scale <= MAX_PRECISION).then_some(Self { unscaled, scale })
40 }
41
42 pub fn unscaled(&self) -> i128 {
44 self.unscaled
45 }
46
47 pub fn scale(&self) -> u32 {
50 self.scale
51 }
52
53 pub fn precision(&self) -> u32 {
56 let mut n = self.unscaled.unsigned_abs();
57 if n == 0 {
58 return 1;
59 }
60 let mut digits = 0;
61 while n > 0 {
62 digits += 1;
63 n /= 10;
64 }
65 digits.max(self.scale)
68 }
69
70 pub fn is_zero(&self) -> bool {
72 self.unscaled == 0
73 }
74
75 pub fn parse(s: &str) -> Option<Self> {
85 let s = s.trim();
86 if s.is_empty() {
87 return None;
88 }
89
90 let (mantissa, exp) = match s.find(['e', 'E']) {
92 Some(i) => {
93 let e: i32 = s[i + 1..].parse().ok()?;
94 (&s[..i], e)
95 }
96 None => (s, 0),
97 };
98
99 let (neg, digits) = match mantissa.strip_prefix('-') {
100 Some(rest) => (true, rest),
101 None => (false, mantissa.strip_prefix('+').unwrap_or(mantissa)),
102 };
103
104 let (int_part, frac_part) = match digits.find('.') {
105 Some(i) => (&digits[..i], &digits[i + 1..]),
106 None => (digits, ""),
107 };
108 if int_part.is_empty() && frac_part.is_empty() {
109 return None;
110 }
111 if !int_part.bytes().all(|b| b.is_ascii_digit())
112 || !frac_part.bytes().all(|b| b.is_ascii_digit())
113 {
114 return None;
115 }
116
117 let mut unscaled: i128 = 0;
119 for b in int_part.bytes().chain(frac_part.bytes()) {
120 unscaled = unscaled.checked_mul(10)?.checked_add((b - b'0') as i128)?;
121 }
122
123 let scale = i64::from(frac_part.len() as u32) - i64::from(exp);
124 let (unscaled, scale) = if scale < 0 {
125 let mut u = unscaled;
129 for _ in 0..(-scale) {
130 u = u.checked_mul(10)?;
131 }
132 (u, 0u32)
133 } else {
134 (unscaled, u32::try_from(scale).ok()?)
135 };
136 if scale > MAX_PRECISION {
137 return None;
138 }
139
140 Some(Self {
141 unscaled: if neg { -unscaled } else { unscaled },
142 scale,
143 })
144 }
145
146 pub fn rescale(&self, target: u32) -> Option<Self> {
149 if target > MAX_PRECISION {
150 return None;
151 }
152 match target.cmp(&self.scale) {
153 Ordering::Equal => Some(*self),
154 Ordering::Greater => {
155 let mut u = self.unscaled;
156 for _ in 0..(target - self.scale) {
157 u = u.checked_mul(10)?;
158 }
159 Some(Self {
160 unscaled: u,
161 scale: target,
162 })
163 }
164 Ordering::Less => {
165 let mut u = self.unscaled;
166 for _ in 0..(self.scale - target) {
167 if u % 10 != 0 {
168 return None; }
170 u /= 10;
171 }
172 Some(Self {
173 unscaled: u,
174 scale: target,
175 })
176 }
177 }
178 }
179
180 pub fn cmp_value(&self, other: &Self) -> Ordering {
188 let target = self.scale.max(other.scale);
189 if let (Some(a), Some(b)) = (self.rescale(target), other.rescale(target)) {
190 return a.unscaled.cmp(&b.unscaled);
191 }
192 digitwise_cmp(self, other)
193 }
194}
195
196fn digitwise_cmp(a: &Decimal, b: &Decimal) -> Ordering {
200 let (sa, sb) = (a.unscaled.signum(), b.unscaled.signum());
201 if sa != sb {
202 return sa.cmp(&sb);
203 }
204 let flip = sa < 0;
205 let (ai, af) = split_digits(a);
206 let (bi, bf) = split_digits(b);
207
208 let (ai, bi) = (ai.trim_start_matches('0'), bi.trim_start_matches('0'));
210 let ord = ai
211 .len()
212 .cmp(&bi.len())
213 .then_with(|| ai.cmp(bi))
214 .then_with(|| {
215 let n = af.len().max(bf.len());
217 let pad = |s: &str| {
218 let mut t = s.to_string();
219 t.extend(std::iter::repeat_n('0', n - s.len()));
220 t
221 };
222 pad(&af).cmp(&pad(&bf))
223 });
224 if flip {
225 ord.reverse()
226 } else {
227 ord
228 }
229}
230
231fn split_digits(d: &Decimal) -> (String, String) {
233 let digits = d.unscaled.unsigned_abs().to_string();
234 let scale = d.scale as usize;
235 if digits.len() > scale {
236 let (i, f) = digits.split_at(digits.len() - scale);
237 (i.to_string(), f.to_string())
238 } else {
239 let mut f = "0".repeat(scale - digits.len());
240 f.push_str(&digits);
241 ("0".into(), f)
242 }
243}
244
245impl fmt::Display for Decimal {
246 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
250 let (int, frac) = split_digits(self);
251 if self.unscaled < 0 {
252 f.write_str("-")?;
253 }
254 f.write_str(&int)?;
255 if self.scale > 0 {
256 write!(f, ".{frac}")?;
257 }
258 Ok(())
259 }
260}
261
262impl PartialOrd for Decimal {
265 fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
266 Some(self.cmp_value(other))
267 }
268}
269
270impl Serialize for Decimal {
274 fn serialize<S: Serializer>(&self, s: S) -> Result<S::Ok, S::Error> {
275 s.collect_str(self)
276 }
277}
278
279impl<'de> Deserialize<'de> for Decimal {
280 fn deserialize<D: Deserializer<'de>>(d: D) -> Result<Self, D::Error> {
281 use serde::de::Error;
282 let raw = serde_json::Value::deserialize(d)?;
283 let text = match &raw {
284 serde_json::Value::String(s) => s.clone(),
285 serde_json::Value::Number(n) if n.is_i64() || n.is_u64() => n.to_string(),
294 serde_json::Value::Number(n) => {
295 return Err(D::Error::custom(format!(
296 "decimal {n} arrived as a JSON number, which is parsed as f64 and has \
297 already lost precision; send it as a string"
298 )))
299 }
300 other => {
301 return Err(D::Error::custom(format!(
302 "decimal must be a string: {other}"
303 )))
304 }
305 };
306 Decimal::parse(&text)
307 .ok_or_else(|| D::Error::custom(format!("not an exact decimal: {text}")))
308 }
309}
310
311#[cfg(test)]
312mod tests {
313 use super::*;
314
315 #[test]
316 fn money_that_f64_corrupts_survives_to_the_digit() {
317 let s = "12345678901234567.89";
319 let d = Decimal::parse(s).unwrap();
320 assert_eq!(d.to_string(), s);
321 assert_eq!(d.unscaled(), 1_234_567_890_123_456_789);
322 assert_eq!(d.scale(), 2);
323 assert_ne!(s.parse::<f64>().unwrap().to_string(), s);
325 }
326
327 #[test]
328 fn scale_is_preserved_not_trimmed() {
329 let d = Decimal::parse("12.3400").unwrap();
330 assert_eq!(d.scale(), 4);
331 assert_eq!(d.to_string(), "12.3400");
332 let e = Decimal::parse("12.34").unwrap();
334 assert_eq!(d.cmp_value(&e), Ordering::Equal);
335 assert_ne!(d, e);
336 }
337
338 #[test]
339 fn parses_sign_fraction_and_exponent_exactly() {
340 assert_eq!(Decimal::parse("-0.004").unwrap().to_string(), "-0.004");
341 assert_eq!(Decimal::parse("1.5e3").unwrap().to_string(), "1500");
342 assert_eq!(Decimal::parse("15e-3").unwrap().to_string(), "0.015");
343 assert_eq!(Decimal::parse(".5").unwrap().to_string(), "0.5");
344 assert_eq!(Decimal::parse("+7").unwrap().to_string(), "7");
345 }
346
347 #[test]
348 fn rejects_what_it_cannot_represent_exactly() {
349 assert!(Decimal::parse("abc").is_none());
350 assert!(Decimal::parse("1.2.3").is_none());
351 assert!(Decimal::parse("").is_none());
352 assert!(Decimal::parse("NaN").is_none());
353 assert!(Decimal::parse(&"9".repeat(40)).is_none());
355 }
356
357 #[test]
358 fn ordering_is_exact_across_scales_and_signs() {
359 let ordered = [
360 "-100", "-2.5", "-0.001", "0", "0.001", "0.0010", "1.5", "1.50", "2.5", "100",
361 ];
362 for w in ordered.windows(2) {
363 let (a, b) = (Decimal::parse(w[0]).unwrap(), Decimal::parse(w[1]).unwrap());
364 assert!(
365 a.cmp_value(&b) != Ordering::Greater,
366 "{} should sort <= {}",
367 w[0],
368 w[1]
369 );
370 }
371 let a = Decimal::parse("100000000000000000.01").unwrap();
373 let b = Decimal::parse("100000000000000000.02").unwrap();
374 assert_eq!(a.cmp_value(&b), Ordering::Less);
375 }
376
377 #[test]
378 fn digitwise_fallback_matches_aligned_compare() {
379 let a = Decimal::from_parts(i128::MAX, 0).unwrap();
382 let b = Decimal::from_parts(1, 30).unwrap();
383 assert_eq!(a.cmp_value(&b), Ordering::Greater);
384 assert_eq!(b.cmp_value(&a), Ordering::Less);
385 let c = Decimal::from_parts(-1, 30).unwrap();
386 assert_eq!(c.cmp_value(&b), Ordering::Less);
387 }
388
389 #[test]
390 fn precision_counts_significant_digits() {
391 assert_eq!(Decimal::parse("0").unwrap().precision(), 1);
392 assert_eq!(Decimal::parse("123.45").unwrap().precision(), 5);
393 assert_eq!(Decimal::parse("0.004").unwrap().precision(), 3);
394 }
395
396 #[test]
397 fn json_is_a_string_and_round_trips() {
398 let d = Decimal::parse("12345678901234567.89").unwrap();
399 let j = serde_json::to_string(&d).unwrap();
400 assert_eq!(j, "\"12345678901234567.89\"");
401 assert_eq!(serde_json::from_str::<Decimal>(&j).unwrap(), d);
402 assert_eq!(
404 serde_json::from_str::<Decimal>("5").unwrap().to_string(),
405 "5"
406 );
407 let err = serde_json::from_str::<Decimal>("1.10")
411 .unwrap_err()
412 .to_string();
413 assert!(err.contains("send it as a string"), "{err}");
414 }
415}