use std::collections::BTreeSet;
use std::fmt::{self, Display};
use std::str::FromStr;
use serde::{de, Deserialize, Deserializer, Serializer};
pub mod comma_separated {
use super::*;
pub fn serialize<T, S>(set: &BTreeSet<T>, serializer: S) -> Result<S::Ok, S::Error>
where
T: Display,
S: Serializer,
{
let s = set
.iter()
.map(|v| v.to_string())
.collect::<Vec<_>>()
.join(",");
serializer.serialize_str(&s)
}
pub fn deserialize<'de, T, D>(deserializer: D) -> Result<BTreeSet<T>, D::Error>
where
T: FromStr + Ord,
T::Err: Display,
D: Deserializer<'de>,
{
let s = String::deserialize(deserializer)?;
if s.is_empty() {
return Ok(BTreeSet::new());
}
s.split(',')
.map(|part| part.trim().parse().map_err(de::Error::custom))
.collect()
}
}
pub mod display_fromstr {
use super::*;
pub fn serialize<T, S>(value: &T, serializer: S) -> Result<S::Ok, S::Error>
where
T: Display,
S: Serializer,
{
serializer.serialize_str(&value.to_string())
}
pub fn deserialize<'de, T, D>(deserializer: D) -> Result<T, D::Error>
where
T: FromStr,
T::Err: Display,
D: Deserializer<'de>,
{
let s = String::deserialize(deserializer)?;
s.parse().map_err(de::Error::custom)
}
}
pub mod default_on_error {
use super::*;
pub fn deserialize<'de, T, D>(deserializer: D) -> Result<Option<T>, D::Error>
where
T: Deserialize<'de>,
D: Deserializer<'de>,
{
Ok(T::deserialize(deserializer).ok())
}
}
pub mod maybe_decimal {
use super::*;
use rust_decimal::Decimal;
pub fn deserialize<'de, D>(deserializer: D) -> Result<Option<Decimal>, D::Error>
where
D: Deserializer<'de>,
{
struct MaybeDecimalVisitor;
impl<'de> de::Visitor<'de> for MaybeDecimalVisitor {
type Value = Option<Decimal>;
fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
formatter.write_str("a decimal string or false")
}
fn visit_bool<E>(self, v: bool) -> Result<Self::Value, E>
where
E: de::Error,
{
if v {
Err(de::Error::custom("expected false or decimal string"))
} else {
Ok(None)
}
}
fn visit_str<E>(self, v: &str) -> Result<Self::Value, E>
where
E: de::Error,
{
v.parse().map(Some).map_err(de::Error::custom)
}
fn visit_string<E>(self, v: String) -> Result<Self::Value, E>
where
E: de::Error,
{
self.visit_str(&v)
}
fn visit_none<E>(self) -> Result<Self::Value, E>
where
E: de::Error,
{
Ok(None)
}
fn visit_unit<E>(self) -> Result<Self::Value, E>
where
E: de::Error,
{
Ok(None)
}
}
deserializer.deserialize_any(MaybeDecimalVisitor)
}
}
pub mod empty_string_as_none {
use super::*;
pub fn deserialize<'de, D>(deserializer: D) -> Result<Option<String>, D::Error>
where
D: Deserializer<'de>,
{
let s = Option::<String>::deserialize(deserializer)?;
Ok(s.filter(|s| !s.is_empty()))
}
}
pub mod optional_comma_separated {
use super::*;
pub fn serialize<T, S>(set: &Option<BTreeSet<T>>, serializer: S) -> Result<S::Ok, S::Error>
where
T: Display,
S: Serializer,
{
match set {
Some(set) => comma_separated::serialize(set, serializer),
None => serializer.serialize_none(),
}
}
pub fn deserialize<'de, T, D>(deserializer: D) -> Result<Option<BTreeSet<T>>, D::Error>
where
T: FromStr + Ord,
T::Err: Display,
D: Deserializer<'de>,
{
let opt: Option<String> = Option::deserialize(deserializer)?;
match opt {
Some(s) if !s.is_empty() => {
let set: Result<BTreeSet<T>, _> = s
.split(',')
.map(|part| part.trim().parse().map_err(de::Error::custom))
.collect();
set.map(Some)
}
_ => Ok(None),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use rust_decimal::Decimal;
use serde::{Deserialize, Serialize};
use std::str::FromStr;
#[test]
fn test_comma_separated_serialize() {
#[derive(Serialize)]
struct Test {
#[serde(with = "comma_separated")]
flags: BTreeSet<String>,
}
let test = Test {
flags: ["a", "b", "c"].iter().map(|s| s.to_string()).collect(),
};
let json = serde_json::to_string(&test).unwrap();
assert_eq!(json, r#"{"flags":"a,b,c"}"#);
}
#[test]
fn test_comma_separated_deserialize() {
#[derive(Deserialize, Debug, PartialEq)]
struct Test {
#[serde(with = "comma_separated")]
flags: BTreeSet<String>,
}
let json = r#"{"flags":"a,b,c"}"#;
let test: Test = serde_json::from_str(json).unwrap();
assert_eq!(test.flags.len(), 3);
assert!(test.flags.contains("a"));
assert!(test.flags.contains("b"));
assert!(test.flags.contains("c"));
}
#[test]
fn test_comma_separated_empty() {
#[derive(Deserialize, Debug)]
struct Test {
#[serde(with = "comma_separated")]
flags: BTreeSet<String>,
}
let json = r#"{"flags":""}"#;
let test: Test = serde_json::from_str(json).unwrap();
assert!(test.flags.is_empty());
}
#[test]
fn test_display_fromstr_serialize() {
#[derive(Serialize)]
struct Test {
#[serde(with = "display_fromstr")]
validate: bool,
}
let test = Test { validate: true };
let json = serde_json::to_string(&test).unwrap();
assert_eq!(json, r#"{"validate":"true"}"#);
}
#[test]
fn test_display_fromstr_deserialize() {
#[derive(Deserialize, Debug)]
struct Test {
#[serde(with = "display_fromstr")]
validate: bool,
}
let json = r#"{"validate":"true"}"#;
let test: Test = serde_json::from_str(json).unwrap();
assert!(test.validate);
let json = r#"{"validate":"false"}"#;
let test: Test = serde_json::from_str(json).unwrap();
assert!(!test.validate);
}
#[test]
fn test_default_on_error_invalid() {
#[derive(Deserialize, Debug)]
struct Test {
#[serde(deserialize_with = "default_on_error::deserialize", default)]
value: Option<i32>,
}
let json = r#"{"value":"not_a_number"}"#;
let test: Test = serde_json::from_str(json).unwrap();
assert!(test.value.is_none());
}
#[test]
fn test_default_on_error_valid() {
#[derive(Deserialize, Debug)]
struct Test {
#[serde(deserialize_with = "default_on_error::deserialize", default)]
value: Option<i32>,
}
let json = r#"{"value":42}"#;
let test: Test = serde_json::from_str(json).unwrap();
assert_eq!(test.value, Some(42));
}
#[test]
fn test_maybe_decimal_false() {
#[derive(Deserialize, Debug)]
struct Test {
#[serde(deserialize_with = "maybe_decimal::deserialize", default)]
limit: Option<Decimal>,
}
let json = r#"{"limit":false}"#;
let test: Test = serde_json::from_str(json).unwrap();
assert!(test.limit.is_none());
}
#[test]
fn test_maybe_decimal_string() {
#[derive(Deserialize, Debug)]
struct Test {
#[serde(deserialize_with = "maybe_decimal::deserialize", default)]
limit: Option<Decimal>,
}
let json = r#"{"limit":"100.50"}"#;
let test: Test = serde_json::from_str(json).unwrap();
assert_eq!(test.limit.unwrap(), Decimal::from_str("100.50").unwrap());
}
#[test]
fn test_empty_string_as_none() {
#[derive(Deserialize, Debug)]
struct Test {
#[serde(deserialize_with = "empty_string_as_none::deserialize", default)]
refid: Option<String>,
}
let json = r#"{"refid":""}"#;
let test: Test = serde_json::from_str(json).unwrap();
assert!(test.refid.is_none());
let json = r#"{"refid":"ABC123"}"#;
let test: Test = serde_json::from_str(json).unwrap();
assert_eq!(test.refid.unwrap(), "ABC123");
}
}