1use std::cmp::Ordering;
2use std::fmt;
3use std::str::FromStr;
4
5use serde::{Deserialize, Deserializer, Serialize, Serializer};
6
7use crate::AssetId;
8
9pub const MAX_DECIMAL_SCALE: u8 = 18;
11
12#[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#[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#[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}