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
102#[cfg(test)]
103mod tests {
104 use serde::de::{value::StrDeserializer, IntoDeserializer};
105 use serde_with::DeserializeAs;
106
107 use super::PermissiveBool;
108
109 fn parse_bool(v: bool) -> Result<bool, serde::de::value::Error> {
110 PermissiveBool::deserialize_as(v.into_deserializer())
111 }
112
113 fn parse_str(s: &str) -> Result<bool, serde::de::value::Error> {
114 let de: StrDeserializer<serde::de::value::Error> = s.into_deserializer();
115 PermissiveBool::deserialize_as(de)
116 }
117
118 fn parse_int(v: i64) -> Result<bool, serde::de::value::Error> {
119 PermissiveBool::deserialize_as(v.into_deserializer())
120 }
121
122 fn parse_uint(v: u64) -> Result<bool, serde::de::value::Error> {
123 PermissiveBool::deserialize_as(v.into_deserializer())
124 }
125
126 fn parse_float(v: f64) -> Result<bool, serde::de::value::Error> {
127 PermissiveBool::deserialize_as(v.into_deserializer())
128 }
129
130 #[test]
132 fn native_bool_is_passed_through() {
133 assert!(parse_bool(true).unwrap());
134 assert!(!parse_bool(false).unwrap());
135 }
136
137 #[test]
139 fn str_truthy() {
140 for s in &["1", "t", "T", "true", "True", "tRuE"] {
141 assert!(parse_str(s).unwrap(), "expected {s:?} to be truthy");
142 }
143 }
144
145 #[test]
146 fn str_falsy() {
147 for s in &["0", "f", "F", "false", "False", "fAlSe"] {
148 assert!(!parse_str(s).unwrap(), "expected {s:?} to be falsy");
149 }
150 }
151
152 #[test]
154 fn str_invalid_rejected() {
155 assert!(parse_str("yes").is_err());
156 assert!(parse_str("no").is_err());
157 assert!(parse_str("2").is_err());
158 assert!(parse_str("").is_err());
159 }
160
161 #[test]
163 fn integer_accepts_zero_and_one() {
164 assert!(!parse_int(0).unwrap());
165 assert!(parse_int(1).unwrap());
166 assert!(!parse_uint(0).unwrap());
167 assert!(parse_uint(1).unwrap());
168 }
169
170 #[test]
171 fn integer_rejects_values_other_than_zero_or_one() {
172 for v in [-1, 2, i64::MIN, i64::MAX] {
174 assert!(parse_int(v).is_err(), "expected signed {v} to be rejected");
175 }
176
177 for v in [2u64, u64::MAX] {
179 assert!(parse_uint(v).is_err(), "expected unsigned {v} to be rejected");
180 }
181 }
182
183 #[test]
185 fn float_accepts_zero_and_one() {
186 assert!(!parse_float(0.0).unwrap());
187 assert!(parse_float(1.0).unwrap());
188 }
189
190 #[test]
191 fn float_rejects_values_other_than_zero_or_one() {
192 for v in [0.5, 2.0, -1.0, f64::NAN, f64::INFINITY, f64::NEG_INFINITY] {
195 assert!(parse_float(v).is_err(), "expected float {v} to be rejected");
196 }
197 }
198}