1use crate::{AvroResult, Error, error::Details};
19use num_bigint::{BigInt, Sign};
20use serde::{Deserialize, Serialize, Serializer, de::SeqAccess};
21use std::num::NonZero;
22
23#[derive(Debug, Clone, Eq)]
24pub struct Decimal {
25 value: BigInt,
26 len: NonZero<usize>,
27}
28
29impl Serialize for Decimal {
30 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
31 where
32 S: Serializer,
33 {
34 match self.to_vec() {
35 Ok(ref bytes) => serializer.serialize_bytes(bytes),
36 Err(e) => Err(serde::ser::Error::custom(e)),
37 }
38 }
39}
40impl<'de> Deserialize<'de> for Decimal {
41 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
42 where
43 D: serde::Deserializer<'de>,
44 {
45 struct DecimalVisitor;
46 impl<'de> serde::de::Visitor<'de> for DecimalVisitor {
47 type Value = Decimal;
48
49 fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result {
50 formatter.write_str("a byte slice or seq of bytes")
51 }
52
53 fn visit_bytes<E>(self, v: &[u8]) -> Result<Self::Value, E>
54 where
55 E: serde::de::Error,
56 {
57 Decimal::new(v).map_err(E::custom)
58 }
59 fn visit_seq<A>(self, mut seq: A) -> Result<Self::Value, A::Error>
60 where
61 A: SeqAccess<'de>,
62 {
63 use serde::de::Error;
64 let mut bytes = Vec::new();
65 while let Some(value) = seq.next_element::<u8>()? {
66 bytes.push(value);
67 }
68
69 Decimal::new(bytes).map_err(A::Error::custom)
70 }
71 }
72 deserializer.deserialize_bytes(DecimalVisitor)
73 }
74}
75
76impl PartialEq for Decimal {
79 fn eq(&self, other: &Self) -> bool {
80 self.value == other.value
81 }
82}
83
84impl Decimal {
85 pub fn new(bytes: impl AsRef<[u8]>) -> AvroResult<Self> {
86 let bytes_ref = bytes.as_ref();
87 Ok(Self {
88 value: BigInt::from_signed_bytes_be(bytes_ref),
89 len: NonZero::new(bytes_ref.len()).ok_or(Details::DecimalIsZeroLength)?,
90 })
91 }
92
93 pub(crate) fn len(&self) -> usize {
94 self.len.get()
95 }
96
97 pub(crate) fn to_vec(&self) -> AvroResult<Vec<u8>> {
98 self.to_sign_extended_bytes_with_len(self.len.get())
99 }
100
101 pub(crate) fn to_sign_extended_bytes_with_len(&self, len: usize) -> AvroResult<Vec<u8>> {
102 let sign_byte = 0xFF * u8::from(self.value.sign() == Sign::Minus);
103 let mut decimal_bytes = vec![sign_byte; len];
104 let raw_bytes = self.value.to_signed_bytes_be();
105 let num_raw_bytes = raw_bytes.len();
106 let start_byte_index = len.checked_sub(num_raw_bytes).ok_or(Details::SignExtend {
107 requested: len,
108 needed: num_raw_bytes,
109 })?;
110 decimal_bytes[start_byte_index..].copy_from_slice(&raw_bytes);
111 Ok(decimal_bytes)
112 }
113}
114
115impl From<Decimal> for BigInt {
116 fn from(decimal: Decimal) -> Self {
117 decimal.value
118 }
119}
120
121impl TryFrom<&Decimal> for Vec<u8> {
132 type Error = Error;
133
134 fn try_from(decimal: &Decimal) -> Result<Self, Self::Error> {
135 decimal.to_vec()
136 }
137}
138
139impl TryFrom<Decimal> for Vec<u8> {
150 type Error = Error;
151
152 fn try_from(decimal: Decimal) -> Result<Self, Self::Error> {
153 decimal.to_vec()
154 }
155}
156
157#[cfg(test)]
158mod tests {
159 use super::*;
160 use apache_avro_test_helper::TestResult;
161 use pretty_assertions::assert_eq;
162
163 #[test]
164 fn test_decimal_from_bytes_from_ref_decimal() -> TestResult {
165 let input = vec![1, 24];
166 let d = Decimal::new(&input)?;
167
168 let output = <Vec<u8>>::try_from(&d)?;
169 assert_eq!(output, input);
170
171 Ok(())
172 }
173
174 #[test]
175 fn test_decimal_from_bytes_from_owned_decimal() -> TestResult {
176 let input = vec![1, 24];
177 let d = Decimal::new(&input)?;
178
179 let output = <Vec<u8>>::try_from(d)?;
180 assert_eq!(output, input);
181
182 Ok(())
183 }
184
185 #[test]
186 fn avro_3949_decimal_serde() -> TestResult {
187 let decimal = Decimal::new([1, 2, 3])?;
188
189 let ser = serde_json::to_string(&decimal)?;
190 let de = serde_json::from_str(&ser)?;
191 std::assert_eq!(decimal, de);
192
193 Ok(())
194 }
195}