massa_models/amount.rs
1// Copyright (c) 2022 MASSA LABS <info@massa.net>
2
3use crate::error::ModelsError;
4use massa_serialization::{Deserializer, SerializeError, Serializer};
5use massa_serialization::{U64VarIntDeserializer, U64VarIntSerializer};
6use nom::error::{context, ContextError, ParseError};
7use nom::{IResult, Parser};
8use rust_decimal::prelude::*;
9use std::fmt;
10use std::ops::Bound;
11use std::str::FromStr;
12
13/// Decimals scale for the amount
14pub const AMOUNT_DECIMAL_SCALE: u32 = 9;
15/// Decimals factor for the amount
16pub const AMOUNT_DECIMAL_FACTOR: u64 = 10u64.pow(AMOUNT_DECIMAL_SCALE);
17
18/// A structure representing a decimal Amount of coins with safe operations
19/// this allows ensuring that there is never an uncontrolled overflow or precision loss
20/// while providing a convenient decimal interface for users
21/// The underlying `u64` raw representation if a fixed-point value with factor `AMOUNT_DECIMAL_FACTOR`
22/// The minimal value is 0 and the maximal value is 18446744073.709551615(std::u64::MAX/1e9)
23#[derive(Clone, Copy, PartialEq, Eq, Ord, PartialOrd, Default)]
24pub struct Amount(u64);
25
26impl Amount {
27 /// Minimum amount
28 pub const MIN: Amount = Amount::from_raw(u64::MIN);
29 /// Maximum amount
30 pub const MAX: Amount = Amount::from_raw(u64::MAX);
31
32 /// Create a zero Amount
33 pub const fn zero() -> Self {
34 Self(0)
35 }
36
37 /// Convert to decimal
38 fn to_decimal(self) -> Decimal {
39 Decimal::from_u64(self.0)
40 .unwrap() // will never panic
41 .checked_div(AMOUNT_DECIMAL_FACTOR.into()) // will never panic
42 .unwrap() // will never panic
43 }
44
45 /// Create an Amount from a Decimal
46 fn from_decimal(dec: Decimal) -> Result<Self, ModelsError> {
47 let res = dec
48 .checked_mul(AMOUNT_DECIMAL_FACTOR.into())
49 .ok_or_else(|| ModelsError::AmountParseError("amount is too large".to_string()))?;
50 if res.is_sign_negative() {
51 return Err(ModelsError::AmountParseError(
52 "amounts cannot be strictly negative".to_string(),
53 ));
54 }
55 if !res.fract().is_zero() {
56 return Err(ModelsError::AmountParseError(format!(
57 "amounts cannot be more precise than 1/{}",
58 AMOUNT_DECIMAL_FACTOR
59 )));
60 }
61 let res = res.to_u64().ok_or_else(|| {
62 ModelsError::AmountParseError(
63 "amount is too large to be represented as u64".to_string(),
64 )
65 })?;
66 Ok(Amount(res))
67 }
68
69 /// Create an Amount from the form `mantissa / (10^scale)` in a const way.
70 /// WARNING: Panics on any error.
71 /// Used only for constant initialization.
72 ///
73 /// ```
74 /// # use massa_models::amount::Amount;
75 /// # use std::str::FromStr;
76 /// let amount_1: Amount = Amount::from_str("0.042").unwrap();
77 /// let amount_2: Amount = Amount::const_init(42, 3);
78 /// assert_eq!(amount_1, amount_2);
79 /// let amount_1: Amount = Amount::from_str("1000").unwrap();
80 /// let amount_2: Amount = Amount::const_init(1000, 0);
81 /// assert_eq!(amount_1, amount_2);
82 /// ```
83 pub const fn const_init(mantissa: u64, scale: u32) -> Self {
84 let raw_mantissa = (mantissa as u128) * (AMOUNT_DECIMAL_FACTOR as u128);
85 let scale_factor = match 10u128.checked_pow(scale) {
86 Some(v) => v,
87 None => panic!(),
88 };
89 assert!(raw_mantissa.is_multiple_of(scale_factor));
90 let res = raw_mantissa / scale_factor;
91 assert!(res <= (u64::MAX as u128));
92 Self(res as u64)
93 }
94
95 /// Returns the value in the (mantissa, scale) format where
96 /// amount = mantissa * 10^(-scale)
97 /// ```
98 /// # use massa_models::amount::Amount;
99 /// # use massa_models::amount::AMOUNT_DECIMAL_SCALE;
100 /// # use std::str::FromStr;
101 /// let amount = Amount::from_str("0.123456789").unwrap();
102 /// let (mantissa, scale) = amount.to_mantissa_scale();
103 /// assert_eq!(mantissa, 123456789);
104 /// assert_eq!(scale, AMOUNT_DECIMAL_SCALE);
105 /// ```
106 pub fn to_mantissa_scale(&self) -> (u64, u32) {
107 (self.0, AMOUNT_DECIMAL_SCALE)
108 }
109
110 /// Creates an amount in the format mantissa*10^(-scale).
111 /// ```
112 /// # use massa_models::amount::Amount;
113 /// # use std::str::FromStr;
114 /// let a = Amount::from_mantissa_scale(123, 2).unwrap();
115 /// assert_eq!(a.to_string(), "1.23");
116 /// let a = Amount::from_mantissa_scale(123, 100);
117 /// assert!(a.is_err());
118 /// ```
119 pub fn from_mantissa_scale(mantissa: u64, scale: u32) -> Result<Self, ModelsError> {
120 let res = Decimal::try_from_i128_with_scale(mantissa as i128, scale)
121 .map_err(|err| ModelsError::AmountParseError(err.to_string()))?;
122 Amount::from_decimal(res)
123 }
124
125 /// Obtains the underlying raw `u64` representation
126 /// Warning: do not use this unless you know what you are doing
127 /// because the raw value does not take the `AMOUNT_DECIMAL_FACTOR` into account.
128 pub const fn to_raw(&self) -> u64 {
129 self.0
130 }
131
132 /// constructs an `Amount` from the underlying raw `u64` representation
133 /// Warning: do not use this unless you know what you are doing
134 /// because the raw value does not take the `AMOUNT_DECIMAL_FACTOR` into account
135 /// In most cases, you should be using `Amount::from_str("11.23")`
136 pub const fn from_raw(raw: u64) -> Self {
137 Self(raw)
138 }
139
140 /// safely add self to another amount, saturating the result on overflow
141 #[must_use]
142 pub fn saturating_add(self, amount: Amount) -> Self {
143 Amount(self.0.saturating_add(amount.0))
144 }
145
146 /// safely subtract another amount from self, saturating the result on underflow
147 #[must_use]
148 pub fn saturating_sub(self, amount: Amount) -> Self {
149 Amount(self.0.saturating_sub(amount.0))
150 }
151
152 /// returns true if the amount is zero
153 pub fn is_zero(&self) -> bool {
154 self.0 == 0
155 }
156
157 /// safely subtract another amount from self, returning None on underflow
158 /// ```
159 /// # use massa_models::amount::Amount;
160 /// # use std::str::FromStr;
161 /// let amount_1 : Amount = Amount::from_str("42").unwrap();
162 /// let amount_2 : Amount = Amount::from_str("7").unwrap();
163 /// let res : Amount = amount_1.checked_sub(amount_2).unwrap();
164 /// assert_eq!(res, Amount::from_str("35").unwrap())
165 /// ```
166 pub fn checked_sub(self, amount: Amount) -> Option<Self> {
167 self.0.checked_sub(amount.0).map(Amount)
168 }
169
170 /// safely add self to another amount, returning None on overflow
171 /// ```
172 /// # use massa_models::amount::Amount;
173 /// # use std::str::FromStr;
174 /// let amount_1 : Amount = Amount::from_str("42").unwrap();
175 /// let amount_2 : Amount = Amount::from_str("7").unwrap();
176 /// let res : Amount = amount_1.checked_add(amount_2).unwrap();
177 /// assert_eq!(res, Amount::from_str("49").unwrap())
178 /// ```
179 pub fn checked_add(self, amount: Amount) -> Option<Self> {
180 self.0.checked_add(amount.0).map(Amount)
181 }
182
183 /// safely multiply self with a `u64`, returning None on overflow
184 /// ```
185 /// # use massa_models::amount::Amount;
186 /// # use std::str::FromStr;
187 /// let amount_1 : Amount = Amount::from_str("42").unwrap();
188 /// let res : Amount = amount_1.checked_mul_u64(7).unwrap();
189 /// assert_eq!(res, Amount::from_str("294").unwrap())
190 /// ```
191 pub fn checked_mul_u64(self, factor: u64) -> Option<Self> {
192 self.0.checked_mul(factor).map(Amount)
193 }
194
195 /// safely multiply self with a `u64`, saturating the result on overflow
196 /// ```
197 /// # use massa_models::amount::Amount;
198 /// # use std::str::FromStr;
199 /// let amount_1 : Amount = Amount::from_str("42").unwrap();
200 /// let res : Amount = amount_1.saturating_mul_u64(7);
201 /// assert_eq!(res, Amount::from_str("294").unwrap());
202 /// ```
203 #[must_use]
204 pub const fn saturating_mul_u64(self, factor: u64) -> Self {
205 Amount(self.0.saturating_mul(factor))
206 }
207
208 /// safely divide self by a `u64`, returning None if the factor is zero
209 /// ```
210 /// # use massa_models::amount::Amount;
211 /// # use std::str::FromStr;
212 /// let amount_1 : Amount = Amount::from_str("42").unwrap();
213 /// let res : Amount = amount_1.checked_div_u64(7).unwrap();
214 /// assert_eq!(res, Amount::from_str("6").unwrap());
215 /// ```
216 pub fn checked_div_u64(self, factor: u64) -> Option<Self> {
217 self.0.checked_div(factor).map(Amount)
218 }
219
220 /// safely divide self by an amount, returning None if the divisor is zero
221 /// ```
222 /// # use massa_models::amount::Amount;
223 /// # use std::str::FromStr;
224 /// let amount_1 : Amount = Amount::from_str("42").unwrap();
225 /// let amount_2 : Amount = Amount::from_str("7").unwrap();
226 /// let res : u64 = amount_1.checked_div(amount_2).unwrap();
227 /// assert_eq!(res, 6);
228 /// ```
229 pub fn checked_div(self, divisor: Self) -> Option<u64> {
230 self.0.checked_div(divisor.0)
231 }
232
233 /// compute self % divisor, return None if divisor is zero
234 /// ```
235 /// # use massa_models::amount::Amount;
236 /// # use std::str::FromStr;
237 /// let amount_1 : Amount = Amount::from_str("42").unwrap();
238 /// let amount_2 : Amount = Amount::from_str("10").unwrap();
239 /// let res : Amount = amount_1.checked_rem(&amount_2).unwrap();
240 /// assert_eq!(res, Amount::from_str("2").unwrap());
241 /// ```
242 pub fn checked_rem(&self, divisor: &Amount) -> Option<Amount> {
243 Some(Amount(self.0.checked_rem(divisor.0)?))
244 }
245
246 /// compute self % divisor, return None if divisor is zero
247 /// ```
248 /// # use massa_models::amount::Amount;
249 /// # use std::str::FromStr;
250 /// let amount_1 : Amount = Amount::from_str("42").unwrap();
251 /// let res : Amount = amount_1.checked_rem_u64(40000000000).unwrap();
252 /// assert_eq!(res, Amount::from_str("2").unwrap());
253 /// ```
254 pub fn checked_rem_u64(&self, divisor: u64) -> Option<Amount> {
255 Some(Amount(self.0.checked_rem(divisor)?))
256 }
257}
258
259/// display an Amount in decimal string form (like "10.33")
260///
261/// ```
262/// # use massa_models::amount::Amount;
263/// # use std::str::FromStr;
264/// let value = Amount::from_str("11.111").unwrap();
265/// assert_eq!(format!("{}", value), "11.111")
266/// ```
267impl fmt::Display for Amount {
268 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
269 write!(f, "{}", self.to_decimal())
270 }
271}
272
273/// Use display implementation in debug to get the decimal representation
274impl fmt::Debug for Amount {
275 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
276 write!(f, "{}", self)
277 }
278}
279
280/// build an Amount from decimal string form (like "10.33")
281/// note that this will fail if the string format is invalid
282/// or if the conversion would cause an overflow, underflow or precision loss
283///
284/// ```
285/// # use massa_models::amount::Amount;
286/// # use std::str::FromStr;
287/// assert!(Amount::from_str("11.1").is_ok());
288/// assert!(Amount::from_str("11.1111111111111111111111").is_err());
289/// assert!(Amount::from_str("1111111111111111111111").is_err());
290/// assert!(Amount::from_str("-11.1").is_err());
291/// assert!(Amount::from_str("abc").is_err());
292/// ```
293impl FromStr for Amount {
294 type Err = ModelsError;
295
296 fn from_str(str_amount: &str) -> Result<Self, Self::Err> {
297 let res = Decimal::from_str_exact(str_amount)
298 .map_err(|err| ModelsError::AmountParseError(err.to_string()))?;
299 Amount::from_decimal(res)
300 }
301}
302
303/// Serializer for amount
304#[derive(Clone)]
305pub struct AmountSerializer {
306 u64_serializer: U64VarIntSerializer,
307}
308
309impl AmountSerializer {
310 /// Create a new `AmountSerializer`
311 pub fn new() -> Self {
312 Self {
313 u64_serializer: U64VarIntSerializer::new(),
314 }
315 }
316}
317
318impl Default for AmountSerializer {
319 fn default() -> Self {
320 Self::new()
321 }
322}
323
324impl Serializer<Amount> for AmountSerializer {
325 /// ## Example
326 /// ```
327 /// use massa_models::amount::{Amount, AmountSerializer};
328 /// use massa_serialization::Serializer;
329 /// use std::str::FromStr;
330 /// use std::ops::Bound::Included;
331 ///
332 /// let amount = Amount::from_str("11.111").unwrap();
333 /// let serializer = AmountSerializer::new();
334 /// let mut serialized = vec![];
335 /// serializer.serialize(&amount, &mut serialized).unwrap();
336 /// ```
337 fn serialize(&self, value: &Amount, buffer: &mut Vec<u8>) -> Result<(), SerializeError> {
338 self.u64_serializer.serialize(&value.0, buffer)
339 }
340}
341
342/// Deserializer for amount
343#[derive(Clone)]
344pub struct AmountDeserializer {
345 u64_deserializer: U64VarIntDeserializer,
346}
347
348impl AmountDeserializer {
349 /// Create a new `AmountDeserializer`
350 pub fn new(min_amount: Bound<Amount>, max_amount: Bound<Amount>) -> Self {
351 let min = match min_amount {
352 Bound::Included(x) => Bound::Included(x.to_raw()),
353 Bound::Excluded(x) => Bound::Excluded(x.to_raw()),
354 Bound::Unbounded => Bound::Included(0),
355 };
356
357 let max = match max_amount {
358 Bound::Included(x) => Bound::Included(x.to_raw()),
359 Bound::Excluded(x) => Bound::Excluded(x.to_raw()),
360 Bound::Unbounded => Bound::Included(Amount::MAX.to_raw()),
361 };
362
363 Self {
364 u64_deserializer: U64VarIntDeserializer::new(min, max),
365 }
366 }
367}
368
369impl Deserializer<Amount> for AmountDeserializer {
370 /// ## Example
371 /// ```
372 /// use massa_models::amount::{Amount, AmountSerializer, AmountDeserializer};
373 /// use massa_serialization::{Serializer, Deserializer, DeserializeError};
374 /// use std::str::FromStr;
375 /// use std::ops::Bound::Included;
376 ///
377 /// let amount = Amount::from_str("11.111").unwrap();
378 /// let serializer = AmountSerializer::new();
379 /// let deserializer = AmountDeserializer::new(Included(Amount::MIN), Included(Amount::MAX));
380 /// let mut serialized = vec![];
381 /// serializer.serialize(&amount, &mut serialized).unwrap();
382 /// let (rest, amount_deser) = deserializer.deserialize::<DeserializeError>(&serialized).unwrap();
383 /// assert!(rest.is_empty());
384 /// assert_eq!(amount_deser, amount);
385 /// ```
386 fn deserialize<'a, E: ParseError<&'a [u8]> + ContextError<&'a [u8]>>(
387 &self,
388 buffer: &'a [u8],
389 ) -> IResult<&'a [u8], Amount, E> {
390 context("Failed Amount deserialization", |input| {
391 self.u64_deserializer.deserialize(input)
392 })
393 .map(Amount::from_raw)
394 .parse(buffer)
395 }
396}
397
398impl<'de> serde::Deserialize<'de> for Amount {
399 fn deserialize<D>(deserializer: D) -> Result<Amount, D::Error>
400 where
401 D: serde::de::Deserializer<'de>,
402 {
403 deserializer.deserialize_str(AmountVisitor)
404 }
405}
406
407struct AmountVisitor;
408
409impl<'de> serde::de::Visitor<'de> for AmountVisitor {
410 type Value = Amount;
411
412 fn visit_str<E>(self, value: &str) -> Result<Amount, E>
413 where
414 E: serde::de::Error,
415 {
416 // The parse error is propagated instead of reflecting the (attacker-controlled,
417 // possibly large) input via `Unexpected::Str(value)`: every message reachable from
418 // `Amount::from_str` is a fixed static string (rust_decimal's parse errors are
419 // `&'static str` literals, and the sign/precision/range checks use constant
420 // messages), so error-path allocation does not scale with the rejected field length.
421 // Matches the `map_err(E::custom)` pattern used by the other string-based
422 // deserializers (Address, Hash, PublicKey, Signature).
423 Amount::from_str(value).map_err(E::custom)
424 }
425
426 fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
427 write!(
428 formatter,
429 "an Amount type representing a fixed-point currency amount"
430 )
431 }
432}
433
434impl serde::Serialize for Amount {
435 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
436 where
437 S: serde::Serializer,
438 {
439 serializer.serialize_str(&self.to_string())
440 }
441}
442
443#[cfg(test)]
444mod tests {
445 use super::Amount;
446 use std::str::FromStr;
447
448 #[test]
449 fn test_valid_amount_still_deserializes() {
450 let amount: Amount = serde_json::from_str("\"12.34\"").unwrap();
451 assert_eq!(amount, Amount::from_str("12.34").unwrap());
452 }
453
454 #[test]
455 fn test_invalid_amount_error_is_bounded_and_explanatory() {
456 // The parse error is propagated (short static description from the parser),
457 // but the rejected Amount string itself must never be embedded in the error
458 // message, to avoid error-path allocation scaling with the input length.
459 let bogus = "z".repeat(4096);
460 let json = format!("\"{}\"", bogus);
461 let err = serde_json::from_str::<Amount>(&json)
462 .expect_err("invalid amount must fail to deserialize");
463 let msg = err.to_string();
464 assert!(
465 !msg.contains('z'),
466 "error message must not reflect the rejected input"
467 );
468 assert!(msg.len() < 128, "error message should stay short: {msg}");
469 assert!(
470 msg.contains("Invalid decimal: unknown character"),
471 "error message should explain the failure: {msg}"
472 );
473 }
474
475 #[test]
476 fn test_amount_error_messages_are_propagated() {
477 // Negative amounts surface the Amount-level explanation.
478 let err = serde_json::from_str::<Amount>("\"-1.5\"")
479 .expect_err("negative amount must fail to deserialize")
480 .to_string();
481 assert!(
482 err.contains("amounts cannot be strictly negative"),
483 "unexpected error: {err}"
484 );
485
486 // Over-precise amounts surface the Amount-level explanation.
487 let err = serde_json::from_str::<Amount>("\"1.1234567891\"")
488 .expect_err("over-precise amount must fail to deserialize")
489 .to_string();
490 assert!(
491 err.contains("amounts cannot be more precise than"),
492 "unexpected error: {err}"
493 );
494
495 // Syntactically invalid decimals surface the rust_decimal explanation.
496 let err = serde_json::from_str::<Amount>("\"1.2.3\"")
497 .expect_err("invalid decimal must fail to deserialize")
498 .to_string();
499 assert!(
500 err.contains("Invalid decimal: two decimal points"),
501 "unexpected error: {err}"
502 );
503 }
504}