Skip to main content

apache_avro/
decimal.rs

1// Licensed to the Apache Software Foundation (ASF) under one
2// or more contributor license agreements.  See the NOTICE file
3// distributed with this work for additional information
4// regarding copyright ownership.  The ASF licenses this file
5// to you under the Apache License, Version 2.0 (the
6// "License"); you may not use this file except in compliance
7// with the License.  You may obtain a copy of the License at
8//
9//   http://www.apache.org/licenses/LICENSE-2.0
10//
11// Unless required by applicable law or agreed to in writing,
12// software distributed under the License is distributed on an
13// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14// KIND, either express or implied.  See the License for the
15// specific language governing permissions and limitations
16// under the License.
17
18use 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
76// We only care about value equality, not byte length. Can two equal `BigInt`s have two different
77// byte lengths?
78impl 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
121/// Gets the internal byte array representation of a referenced decimal.
122/// Usage:
123/// ```
124/// # use std::convert::TryFrom;
125/// # use apache_avro::{Decimal, Error};
126/// #
127/// let decimal = Decimal::new([1, 24])?;
128/// let maybe_bytes = <Vec<u8>>::try_from(&decimal);
129/// # Ok::<(), Error>(())
130/// ```
131impl 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
139/// Gets the internal byte array representation of an owned decimal.
140/// Usage:
141/// ```
142/// # use std::convert::TryFrom;
143/// # use apache_avro::{Decimal, Error};
144/// #
145/// let decimal = Decimal::new([1, 24])?;
146/// let maybe_bytes = <Vec<u8>>::try_from(decimal);
147/// # Ok::<(), Error>(())
148/// ```
149impl 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}