Skip to main content

salvo_oapi/openapi/schema/
number.rs

1use serde::de::{self, Visitor};
2use serde::{Deserialize, Deserializer, Serialize, Serializer};
3
4/// Represents a numeric value in OpenAPI schema validation fields.
5///
6/// This type preserves the distinction between integers and floating-point numbers,
7/// ensuring that integer values like `1` serialize as `1` rather than `1.0` in JSON output.
8///
9/// # Examples
10///
11/// ```
12/// # use salvo_oapi::schema::Number;
13/// let int_val: Number = 42.into();
14/// let float_val: Number = 3.14.into();
15///
16/// assert_eq!(serde_json::to_string(&int_val).unwrap(), "42");
17/// assert_eq!(serde_json::to_string(&float_val).unwrap(), "3.14");
18/// ```
19#[derive(Clone, Debug)]
20pub enum Number {
21    /// Signed integer value e.g. `1` or `-2`.
22    Int(isize),
23    /// Unsigned integer value e.g. `0`.
24    UInt(usize),
25    /// Floating point number e.g. `1.34`.
26    Float(f64),
27}
28
29impl Eq for Number {}
30
31impl PartialEq for Number {
32    fn eq(&self, other: &Self) -> bool {
33        match (self, other) {
34            (Self::Int(left), Self::Int(right)) => left == right,
35            (Self::UInt(left), Self::UInt(right)) => left == right,
36            // A float is equal when numerically equal (so `-0.0 == 0.0`) or when the
37            // bit patterns match, so `NaN == NaN`. Without the bit check, `Eq` would
38            // be unsound: a plain `f64 == f64` makes `Number::Float(NaN)` unequal to
39            // itself, violating the reflexivity `Eq` promises.
40            (Self::Float(left), Self::Float(right)) => {
41                left == right || left.to_bits() == right.to_bits()
42            }
43            _ => false,
44        }
45    }
46}
47
48impl Serialize for Number {
49    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
50    where
51        S: Serializer,
52    {
53        match self {
54            Self::Int(value) => serializer.serialize_i64(*value as i64),
55            Self::UInt(value) => serializer.serialize_u64(*value as u64),
56            Self::Float(value) => {
57                // Serialize whole floats as integers to avoid trailing `.0`
58                if value.fract() == 0.0 && value.is_finite() {
59                    if *value < 0.0 {
60                        serializer.serialize_i64(*value as i64)
61                    } else {
62                        serializer.serialize_u64(*value as u64)
63                    }
64                } else {
65                    serializer.serialize_f64(*value)
66                }
67            }
68        }
69    }
70}
71
72struct NumberVisitor;
73
74impl<'de> Visitor<'de> for NumberVisitor {
75    type Value = Number;
76
77    fn expecting(&self, formatter: &mut std::fmt::Formatter) -> std::fmt::Result {
78        formatter.write_str("a number (integer or float)")
79    }
80
81    fn visit_i64<E>(self, v: i64) -> Result<Self::Value, E>
82    where
83        E: de::Error,
84    {
85        Ok(Number::Int(v as isize))
86    }
87
88    fn visit_u64<E>(self, v: u64) -> Result<Self::Value, E>
89    where
90        E: de::Error,
91    {
92        Ok(Number::UInt(v as usize))
93    }
94
95    fn visit_f64<E>(self, v: f64) -> Result<Self::Value, E>
96    where
97        E: de::Error,
98    {
99        Ok(Number::Float(v))
100    }
101}
102
103impl<'de> Deserialize<'de> for Number {
104    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
105    where
106        D: Deserializer<'de>,
107    {
108        deserializer.deserialize_any(NumberVisitor)
109    }
110}
111
112macro_rules! impl_from_for_number {
113    ( $( $ty:ident => $pat:ident $( as $as:ident )? ),* ) => {
114        $(
115        impl From<$ty> for Number {
116            fn from(value: $ty) -> Self {
117                Self::$pat(value $( as $as )?)
118            }
119        }
120        )*
121    };
122}
123
124#[rustfmt::skip]
125impl_from_for_number!(
126    f32 => Float as f64, f64 => Float,
127    i8 => Int as isize, i16 => Int as isize, i32 => Int as isize, i64 => Int as isize,
128    u8 => UInt as usize, u16 => UInt as usize, u32 => UInt as usize, u64 => UInt as usize,
129    isize => Int, usize => UInt
130);
131
132#[cfg(test)]
133mod tests {
134    use super::*;
135
136    #[test]
137    fn test_serialize_int() {
138        let n = Number::Int(42);
139        assert_eq!(serde_json::to_string(&n).unwrap(), "42");
140    }
141
142    #[test]
143    fn test_serialize_negative_int() {
144        let n = Number::Int(-5);
145        assert_eq!(serde_json::to_string(&n).unwrap(), "-5");
146    }
147
148    #[test]
149    fn test_serialize_uint() {
150        let n = Number::UInt(100);
151        assert_eq!(serde_json::to_string(&n).unwrap(), "100");
152    }
153
154    #[test]
155    #[allow(clippy::approx_constant)]
156    fn test_serialize_float() {
157        let n = Number::Float(3.14);
158        assert_eq!(serde_json::to_string(&n).unwrap(), "3.14");
159    }
160
161    #[test]
162    fn test_serialize_whole_float_as_integer() {
163        let n = Number::Float(10.0);
164        assert_eq!(serde_json::to_string(&n).unwrap(), "10");
165    }
166
167    #[test]
168    fn test_serialize_negative_whole_float() {
169        let n = Number::Float(-3.0);
170        assert_eq!(serde_json::to_string(&n).unwrap(), "-3");
171    }
172
173    #[test]
174    fn test_from_i32() {
175        let n: Number = 42i32.into();
176        assert_eq!(n, Number::Int(42));
177    }
178
179    #[test]
180    fn test_from_u64() {
181        let n: Number = 100u64.into();
182        assert_eq!(n, Number::UInt(100));
183    }
184
185    #[test]
186    fn test_from_f64() {
187        let n: Number = 2.5f64.into();
188        assert_eq!(n, Number::Float(2.5));
189    }
190
191    #[test]
192    fn test_deserialize_int() {
193        let n: Number = serde_json::from_str("42").unwrap();
194        assert_eq!(n, Number::UInt(42));
195    }
196
197    #[test]
198    fn test_deserialize_negative_int() {
199        let n: Number = serde_json::from_str("-5").unwrap();
200        assert_eq!(n, Number::Int(-5));
201    }
202
203    #[test]
204    #[allow(clippy::approx_constant)]
205    fn test_deserialize_float() {
206        let n: Number = serde_json::from_str("3.14").unwrap();
207        assert_eq!(n, Number::Float(3.14));
208    }
209
210    #[test]
211    fn test_equality() {
212        assert_eq!(Number::Int(1), Number::Int(1));
213        assert_eq!(Number::UInt(1), Number::UInt(1));
214        assert_eq!(Number::Float(1.5), Number::Float(1.5));
215        assert_ne!(Number::Int(1), Number::UInt(1));
216    }
217
218    #[test]
219    fn test_eq_is_reflexive_for_nan() {
220        // `Eq` requires `a == a`; the bit-pattern comparison must hold for NaN even
221        // though `f64::NAN != f64::NAN`.
222        let nan = Number::Float(f64::NAN);
223        assert_eq!(nan, nan);
224        // `-0.0` and `0.0` remain equal (numerically equal).
225        assert_eq!(Number::Float(-0.0), Number::Float(0.0));
226    }
227}