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}