1use std::fmt;
7
8use serde::{
9 de::{Error, Unexpected},
10 Deserializer,
11};
12use serde_with::DeserializeAs;
13
14pub 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 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
102pub 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 #[test]
161 fn native_bool_is_passed_through() {
162 assert!(parse_bool(true).unwrap());
163 assert!(!parse_bool(false).unwrap());
164 }
165
166 #[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 #[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 #[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 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 for v in [2u64, u64::MAX] {
208 assert!(parse_uint(v).is_err(), "expected unsigned {v} to be rejected");
209 }
210 }
211
212 #[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 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}