Skip to main content

qs_instruments/
decimal.rs

1use std::cmp::Ordering;
2use std::fmt;
3use std::str::FromStr;
4
5use serde::{Deserialize, Deserializer, Serialize, Serializer};
6
7use crate::AssetId;
8
9/// Maximum supported number of decimal fractional digits.
10pub const MAX_DECIMAL_SCALE: u8 = 18;
11
12/// An exact checked base-10 decimal backed by an `i128` coefficient.
13#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash)]
14pub struct Decimal {
15    coefficient: i128,
16    scale: u8,
17}
18
19impl Decimal {
20    pub const ZERO: Self = Self {
21        coefficient: 0,
22        scale: 0,
23    };
24
25    pub fn new(coefficient: i128, scale: u8) -> Result<Self, DecimalError> {
26        if scale > MAX_DECIMAL_SCALE {
27            return Err(DecimalError::ScaleTooLarge(scale));
28        }
29        Ok(Self::normalize(coefficient, scale))
30    }
31
32    pub const fn coefficient(self) -> i128 {
33        self.coefficient
34    }
35
36    pub const fn scale(self) -> u8 {
37        self.scale
38    }
39
40    pub const fn is_zero(self) -> bool {
41        self.coefficient == 0
42    }
43
44    pub const fn is_positive(self) -> bool {
45        self.coefficient > 0
46    }
47
48    pub const fn is_negative(self) -> bool {
49        self.coefficient < 0
50    }
51
52    pub fn checked_add(self, rhs: Self) -> Result<Self, DecimalError> {
53        let (left, right, scale) = align(self, rhs)?;
54        let coefficient = left.checked_add(right).ok_or(DecimalError::Overflow)?;
55        Self::new(coefficient, scale)
56    }
57
58    pub fn checked_sub(self, rhs: Self) -> Result<Self, DecimalError> {
59        let (left, right, scale) = align(self, rhs)?;
60        let coefficient = left.checked_sub(right).ok_or(DecimalError::Overflow)?;
61        Self::new(coefficient, scale)
62    }
63
64    pub fn checked_mul(self, rhs: Self) -> Result<Self, DecimalError> {
65        let coefficient = self
66            .coefficient
67            .checked_mul(rhs.coefficient)
68            .ok_or(DecimalError::Overflow)?;
69        let scale = self
70            .scale
71            .checked_add(rhs.scale)
72            .ok_or(DecimalError::Overflow)?;
73        let normalized = Self::normalize(coefficient, scale);
74        if normalized.scale > MAX_DECIMAL_SCALE {
75            return Err(DecimalError::ScaleTooLarge(normalized.scale));
76        }
77        Ok(normalized)
78    }
79
80    pub fn checked_from_f64(value: f64) -> Result<Self, DecimalError> {
81        if !value.is_finite() {
82            return Err(DecimalError::NonFiniteFloat);
83        }
84        value.to_string().parse()
85    }
86
87    pub fn checked_rescale(self, scale: u8) -> Result<Self, DecimalError> {
88        if scale > MAX_DECIMAL_SCALE {
89            return Err(DecimalError::ScaleTooLarge(scale));
90        }
91        if scale == self.scale {
92            return Ok(self);
93        }
94        if scale > self.scale {
95            let factor = power_of_ten(scale - self.scale)?;
96            let coefficient = self
97                .coefficient
98                .checked_mul(factor)
99                .ok_or(DecimalError::Overflow)?;
100            return Self::new(coefficient, scale);
101        }
102
103        let factor = power_of_ten(self.scale - scale)?;
104        if self.coefficient % factor != 0 {
105            return Err(DecimalError::InexactRescale {
106                from: self.scale,
107                to: scale,
108            });
109        }
110        Self::new(self.coefficient / factor, scale)
111    }
112
113    pub(crate) fn aligned_coefficients(self, rhs: Self) -> Result<(i128, i128, u8), DecimalError> {
114        align(self, rhs)
115    }
116
117    fn normalize(mut coefficient: i128, mut scale: u8) -> Self {
118        if coefficient == 0 {
119            return Self::ZERO;
120        }
121        while scale > 0 && coefficient % 10 == 0 {
122            coefficient /= 10;
123            scale -= 1;
124        }
125        Self { coefficient, scale }
126    }
127}
128
129fn align(left: Decimal, right: Decimal) -> Result<(i128, i128, u8), DecimalError> {
130    let scale = left.scale.max(right.scale);
131    let left_factor = power_of_ten(scale - left.scale)?;
132    let right_factor = power_of_ten(scale - right.scale)?;
133    let left = left
134        .coefficient
135        .checked_mul(left_factor)
136        .ok_or(DecimalError::Overflow)?;
137    let right = right
138        .coefficient
139        .checked_mul(right_factor)
140        .ok_or(DecimalError::Overflow)?;
141    Ok((left, right, scale))
142}
143
144fn power_of_ten(power: u8) -> Result<i128, DecimalError> {
145    10_i128
146        .checked_pow(u32::from(power))
147        .ok_or(DecimalError::Overflow)
148}
149
150fn compare_magnitude(left: Decimal, right: Decimal) -> Ordering {
151    let left_digits = left.coefficient.unsigned_abs().to_string();
152    let right_digits = right.coefficient.unsigned_abs().to_string();
153    let left_exponent = left_digits.len() as i32 - i32::from(left.scale);
154    let right_exponent = right_digits.len() as i32 - i32::from(right.scale);
155    match left_exponent.cmp(&right_exponent) {
156        Ordering::Equal => {
157            let length = left_digits.len().max(right_digits.len());
158            left_digits
159                .bytes()
160                .chain(std::iter::repeat(b'0'))
161                .take(length)
162                .cmp(
163                    right_digits
164                        .bytes()
165                        .chain(std::iter::repeat(b'0'))
166                        .take(length),
167                )
168        }
169        ordering => ordering,
170    }
171}
172
173impl Ord for Decimal {
174    fn cmp(&self, other: &Self) -> Ordering {
175        match (self.coefficient.signum(), other.coefficient.signum()) {
176            (left, right) if left != right => left.cmp(&right),
177            (-1, -1) => compare_magnitude(*other, *self),
178            (0, 0) => Ordering::Equal,
179            _ => compare_magnitude(*self, *other),
180        }
181    }
182}
183
184impl PartialOrd for Decimal {
185    fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
186        Some(self.cmp(other))
187    }
188}
189
190impl fmt::Display for Decimal {
191    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
192        if self.scale == 0 {
193            return write!(f, "{}", self.coefficient);
194        }
195
196        let negative = self.coefficient < 0;
197        let digits = self.coefficient.unsigned_abs().to_string();
198        let scale = usize::from(self.scale);
199        if negative {
200            f.write_str("-")?;
201        }
202        if digits.len() <= scale {
203            f.write_str("0.")?;
204            for _ in 0..(scale - digits.len()) {
205                f.write_str("0")?;
206            }
207            f.write_str(&digits)
208        } else {
209            let split = digits.len() - scale;
210            write!(f, "{}.{}", &digits[..split], &digits[split..])
211        }
212    }
213}
214
215impl FromStr for Decimal {
216    type Err = DecimalError;
217
218    fn from_str(value: &str) -> Result<Self, Self::Err> {
219        if value.is_empty() || value.trim() != value {
220            return Err(DecimalError::InvalidFormat);
221        }
222        if value.starts_with('+') || value.contains(['e', 'E']) {
223            return Err(DecimalError::InvalidFormat);
224        }
225
226        let negative = value.starts_with('-');
227        let unsigned = value.strip_prefix('-').unwrap_or(value);
228        let mut components = unsigned.split('.');
229        let integer = components.next().ok_or(DecimalError::InvalidFormat)?;
230        let fractional = components.next();
231        if components.next().is_some()
232            || integer.is_empty()
233            || !integer.bytes().all(|byte| byte.is_ascii_digit())
234            || fractional.is_some_and(|part| {
235                part.is_empty() || !part.bytes().all(|byte| byte.is_ascii_digit())
236            })
237        {
238            return Err(DecimalError::InvalidFormat);
239        }
240
241        let fractional = fractional.unwrap_or("");
242        let scale = u8::try_from(fractional.len()).map_err(|_| DecimalError::InvalidFormat)?;
243        if scale > MAX_DECIMAL_SCALE {
244            return Err(DecimalError::ScaleTooLarge(scale));
245        }
246        let digits = format!("{integer}{fractional}");
247        let magnitude = digits.parse::<u128>().map_err(|_| DecimalError::Overflow)?;
248        let coefficient = if negative {
249            if magnitude == 0 {
250                return Err(DecimalError::NegativeZero);
251            }
252            if magnitude == i128::MAX as u128 + 1 {
253                i128::MIN
254            } else {
255                let coefficient = i128::try_from(magnitude).map_err(|_| DecimalError::Overflow)?;
256                coefficient.checked_neg().ok_or(DecimalError::Overflow)?
257            }
258        } else {
259            i128::try_from(magnitude).map_err(|_| DecimalError::Overflow)?
260        };
261        Self::new(coefficient, scale)
262    }
263}
264
265impl Serialize for Decimal {
266    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
267    where
268        S: Serializer,
269    {
270        serializer.collect_str(self)
271    }
272}
273
274impl<'de> Deserialize<'de> for Decimal {
275    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
276    where
277        D: Deserializer<'de>,
278    {
279        let value = String::deserialize(deserializer)?;
280        value.parse().map_err(serde::de::Error::custom)
281    }
282}
283
284macro_rules! decimal_wrapper {
285    ($name:ident, $predicate:expr, $message:literal) => {
286        #[derive(Clone, Copy, Debug, Eq, PartialEq, Ord, PartialOrd, Hash)]
287        pub struct $name(Decimal);
288
289        impl $name {
290            pub fn new(value: Decimal) -> Result<Self, DecimalError> {
291                if !($predicate)(value) {
292                    return Err(DecimalError::ConstraintViolation($message));
293                }
294                Ok(Self(value))
295            }
296
297            pub const fn get(self) -> Decimal {
298                self.0
299            }
300        }
301
302        impl fmt::Display for $name {
303            fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
304                self.0.fmt(f)
305            }
306        }
307
308        impl FromStr for $name {
309            type Err = DecimalError;
310
311            fn from_str(value: &str) -> Result<Self, Self::Err> {
312                Self::new(value.parse()?)
313            }
314        }
315
316        impl TryFrom<Decimal> for $name {
317            type Error = DecimalError;
318
319            fn try_from(value: Decimal) -> Result<Self, Self::Error> {
320                Self::new(value)
321            }
322        }
323
324        impl From<$name> for Decimal {
325            fn from(value: $name) -> Self {
326                value.get()
327            }
328        }
329
330        impl Serialize for $name {
331            fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
332            where
333                S: Serializer,
334            {
335                self.0.serialize(serializer)
336            }
337        }
338
339        impl<'de> Deserialize<'de> for $name {
340            fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
341            where
342                D: Deserializer<'de>,
343            {
344                Self::new(Decimal::deserialize(deserializer)?).map_err(serde::de::Error::custom)
345            }
346        }
347    };
348}
349
350decimal_wrapper!(
351    PositiveDecimal,
352    |value: Decimal| value.is_positive(),
353    "value must be positive"
354);
355decimal_wrapper!(
356    NonNegativeDecimal,
357    |value: Decimal| !value.is_negative(),
358    "value must be nonnegative"
359);
360decimal_wrapper!(
361    Price,
362    |value: Decimal| value.is_positive(),
363    "price must be positive"
364);
365decimal_wrapper!(
366    Quantity,
367    |value: Decimal| !value.is_negative(),
368    "quantity must be nonnegative"
369);
370
371impl Quantity {
372    pub fn require_positive(self) -> Result<PositiveDecimal, DecimalError> {
373        PositiveDecimal::new(self.get())
374    }
375}
376
377/// An exact amount denominated in one asset.
378#[derive(Clone, Debug, Eq, PartialEq, Serialize, Deserialize)]
379#[serde(deny_unknown_fields)]
380pub struct Money {
381    pub asset: AssetId,
382    pub amount: Decimal,
383}
384
385/// Exact-decimal parsing and arithmetic failures.
386#[derive(Clone, Debug, Eq, PartialEq, thiserror::Error)]
387pub enum DecimalError {
388    #[error("invalid decimal string")]
389    InvalidFormat,
390    #[error("decimal scale {0} exceeds maximum scale {MAX_DECIMAL_SCALE}")]
391    ScaleTooLarge(u8),
392    #[error("decimal arithmetic overflow")]
393    Overflow,
394    #[error("negative zero is not canonical")]
395    NegativeZero,
396    #[error("non-finite floating-point values cannot be converted to Decimal")]
397    NonFiniteFloat,
398    #[error("cannot rescale exactly from scale {from} to scale {to}")]
399    InexactRescale { from: u8, to: u8 },
400    #[error("{0}")]
401    ConstraintViolation(&'static str),
402}