saluki_common/
deser.rs

1//! Deserialization helpers.
2//!
3//! This module provides various helpers for handling the deserialization of common data types in more flexible and
4//! permissive ways. These helpers are designed to be used with the `serde_with` crate.
5
6use std::fmt;
7
8use serde::{
9    de::{Error, Unexpected},
10    Deserializer,
11};
12use serde_with::DeserializeAs;
13
14/// Permissively deserializes a boolean.
15///
16/// This helper module allows deserializing a `bool` from a number of possible data types:
17///
18/// - `true` or `false` as a native boolean
19/// - `1` or `0` as an integer (signed, unsigned, or floating point)
20/// - `"true"` or `"false"` as a string, case insensitive
21/// - `"1"`, `"t"`, or `"T"` as truthy strings; `"0"`, `"f"`, or `"F"` as falsy strings
22pub struct PermissiveBool;
23
24impl<'de> DeserializeAs<'de, bool> for PermissiveBool {
25    fn deserialize_as<D>(deserializer: D) -> Result<bool, D::Error>
26    where
27        D: Deserializer<'de>,
28    {
29        struct Visitor;
30
31        impl<'vde> serde::de::Visitor<'vde> for Visitor {
32            type Value = bool;
33
34            fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
35                formatter.write_str("a boolean, string, integer, or floating-point number")
36            }
37
38            fn visit_bool<E>(self, value: bool) -> Result<Self::Value, E>
39            where
40                E: Error,
41            {
42                Ok(value)
43            }
44
45            fn visit_str<E>(self, value: &str) -> Result<Self::Value, E>
46            where
47                E: Error,
48            {
49                // Check short forms first, then fall back to case-insensitive "true"/"false".
50                match value {
51                    "1" | "t" | "T" => Ok(true),
52                    "0" | "f" | "F" => Ok(false),
53                    _ => match value.to_lowercase().as_str() {
54                        "true" => Ok(true),
55                        "false" => Ok(false),
56                        _ => Err(Error::invalid_value(
57                            Unexpected::Str(value),
58                            &"a boolean string (\"true\" or \"false\", case insensitive, or short forms: \"1\", \"t\", \"T\", \"0\", \"f\", \"F\")",
59                        )),
60                    },
61                }
62            }
63
64            fn visit_i64<E>(self, value: i64) -> Result<Self::Value, E>
65            where
66                E: Error,
67            {
68                match value {
69                    0 => Ok(false),
70                    1 => Ok(true),
71                    _ => Err(Error::invalid_value(Unexpected::Signed(value), &"0 or 1")),
72                }
73            }
74
75            fn visit_u64<E>(self, value: u64) -> Result<Self::Value, E>
76            where
77                E: Error,
78            {
79                match value {
80                    0 => Ok(false),
81                    1 => Ok(true),
82                    _ => Err(Error::invalid_value(Unexpected::Unsigned(value), &"0 or 1")),
83                }
84            }
85
86            fn visit_f64<E>(self, value: f64) -> Result<Self::Value, E>
87            where
88                E: Error,
89            {
90                match value {
91                    0.0 => Ok(false),
92                    1.0 => Ok(true),
93                    _ => Err(Error::invalid_value(Unexpected::Float(value), &"0.0 or 1.0")),
94                }
95            }
96        }
97
98        deserializer.deserialize_any(Visitor)
99    }
100}
101
102/// Deserializes an optional string field, returning `None` if a string is found but is empty.
103pub fn empty_string_as_none<'de, D>(deserializer: D) -> Result<Option<String>, D::Error>
104where
105    D: Deserializer<'de>,
106{
107    struct Visitor;
108
109    impl<'vde> serde::de::Visitor<'vde> for Visitor {
110        type Value = Option<String>;
111
112        fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
113            formatter.write_str("a string")
114        }
115
116        fn visit_str<E>(self, value: &str) -> Result<Self::Value, E>
117        where
118            E: Error,
119        {
120            if value.is_empty() {
121                return Ok(None);
122            }
123
124            Ok(Some(value.to_string()))
125        }
126    }
127
128    deserializer.deserialize_any(Visitor)
129}
130
131#[cfg(test)]
132mod tests {
133    use serde::de::{value::StrDeserializer, IntoDeserializer};
134    use serde_with::DeserializeAs;
135
136    use super::PermissiveBool;
137
138    fn parse_bool(v: bool) -> Result<bool, serde::de::value::Error> {
139        PermissiveBool::deserialize_as(v.into_deserializer())
140    }
141
142    fn parse_str(s: &str) -> Result<bool, serde::de::value::Error> {
143        let de: StrDeserializer<serde::de::value::Error> = s.into_deserializer();
144        PermissiveBool::deserialize_as(de)
145    }
146
147    fn parse_int(v: i64) -> Result<bool, serde::de::value::Error> {
148        PermissiveBool::deserialize_as(v.into_deserializer())
149    }
150
151    fn parse_uint(v: u64) -> Result<bool, serde::de::value::Error> {
152        PermissiveBool::deserialize_as(v.into_deserializer())
153    }
154
155    fn parse_float(v: f64) -> Result<bool, serde::de::value::Error> {
156        PermissiveBool::deserialize_as(v.into_deserializer())
157    }
158
159    // Native boolean
160    #[test]
161    fn native_bool_is_passed_through() {
162        assert!(parse_bool(true).unwrap());
163        assert!(!parse_bool(false).unwrap());
164    }
165
166    // String variants
167    #[test]
168    fn str_truthy() {
169        for s in &["1", "t", "T", "true", "True", "tRuE"] {
170            assert!(parse_str(s).unwrap(), "expected {s:?} to be truthy");
171        }
172    }
173
174    #[test]
175    fn str_falsy() {
176        for s in &["0", "f", "F", "false", "False", "fAlSe"] {
177            assert!(!parse_str(s).unwrap(), "expected {s:?} to be falsy");
178        }
179    }
180
181    // Invalid string
182    #[test]
183    fn str_invalid_rejected() {
184        assert!(parse_str("yes").is_err());
185        assert!(parse_str("no").is_err());
186        assert!(parse_str("2").is_err());
187        assert!(parse_str("").is_err());
188    }
189
190    // Integer variants (both signed and unsigned deserializer paths).
191    #[test]
192    fn integer_accepts_zero_and_one() {
193        assert!(!parse_int(0).unwrap());
194        assert!(parse_int(1).unwrap());
195        assert!(!parse_uint(0).unwrap());
196        assert!(parse_uint(1).unwrap());
197    }
198
199    #[test]
200    fn integer_rejects_values_other_than_zero_or_one() {
201        // Signed values outside {0, 1}, including negatives and extremes, are rejected (visit_i64).
202        for v in [-1, 2, i64::MIN, i64::MAX] {
203            assert!(parse_int(v).is_err(), "expected signed {v} to be rejected");
204        }
205
206        // Unsigned values outside {0, 1} are rejected (visit_u64).
207        for v in [2u64, u64::MAX] {
208            assert!(parse_uint(v).is_err(), "expected unsigned {v} to be rejected");
209        }
210    }
211
212    // Floating-point variants (visit_f64 path).
213    #[test]
214    fn float_accepts_zero_and_one() {
215        assert!(!parse_float(0.0).unwrap());
216        assert!(parse_float(1.0).unwrap());
217    }
218
219    #[test]
220    fn float_rejects_values_other_than_zero_or_one() {
221        // Any floating-point value other than exactly 0.0 or 1.0 is rejected, including fractional,
222        // out-of-range, and non-finite values.
223        for v in [0.5, 2.0, -1.0, f64::NAN, f64::INFINITY, f64::NEG_INFINITY] {
224            assert!(parse_float(v).is_err(), "expected float {v} to be rejected");
225        }
226    }
227}