use std::fmt::Display;
use std::str::FromStr;
use base64::engine::general_purpose::{STANDARD, STANDARD_NO_PAD, URL_SAFE, URL_SAFE_NO_PAD};
use base64::Engine as _;
use serde::de::{self, Unexpected, Visitor};
use serde::{Deserialize, Deserializer, Serializer};
#[derive(Debug, Clone, Deserialize)]
#[serde(untagged)]
pub enum EnumRepr {
Name(String),
Number(i32),
}
fn decode_base64<E: de::Error>(s: &str) -> Result<Vec<u8>, E> {
for engine in [&STANDARD, &STANDARD_NO_PAD, &URL_SAFE, &URL_SAFE_NO_PAD] {
if let Ok(bytes) = engine.decode(s) {
return Ok(bytes);
}
}
Err(E::invalid_value(
Unexpected::Str(s),
&"base64-encoded bytes",
))
}
struct IntVisitor<T>(std::marker::PhantomData<T>);
impl<T> Visitor<'_> for IntVisitor<T>
where
T: FromStr + TryFrom<i64> + TryFrom<u64>,
<T as FromStr>::Err: Display,
{
type Value = T;
fn expecting(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("an integer, as a JSON number or a decimal string")
}
fn visit_str<E: de::Error>(self, v: &str) -> Result<T, E> {
v.parse().map_err(|e: <T as FromStr>::Err| E::custom(e))
}
fn visit_i64<E: de::Error>(self, v: i64) -> Result<T, E> {
T::try_from(v).map_err(|_| E::custom(format!("integer {v} out of range")))
}
fn visit_u64<E: de::Error>(self, v: u64) -> Result<T, E> {
T::try_from(v).map_err(|_| E::custom(format!("integer {v} out of range")))
}
}
fn deserialize_int<'de, T, D>(d: D) -> Result<T, D::Error>
where
T: FromStr + TryFrom<i64> + TryFrom<u64>,
<T as FromStr>::Err: Display,
D: Deserializer<'de>,
{
d.deserialize_any(IntVisitor(std::marker::PhantomData))
}
pub mod opt_int {
use super::*;
pub fn serialize<T: Display, S: Serializer>(v: &Option<T>, s: S) -> Result<S::Ok, S::Error> {
match v {
Some(v) => s.serialize_str(&v.to_string()),
None => s.serialize_none(),
}
}
pub fn deserialize<'de, T, D>(d: D) -> Result<Option<T>, D::Error>
where
T: FromStr + TryFrom<i64> + TryFrom<u64>,
<T as FromStr>::Err: Display,
D: Deserializer<'de>,
{
struct OptVisitor<T>(std::marker::PhantomData<T>);
impl<'de, T> Visitor<'de> for OptVisitor<T>
where
T: FromStr + TryFrom<i64> + TryFrom<u64>,
<T as FromStr>::Err: Display,
{
type Value = Option<T>;
fn expecting(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("an optional integer")
}
fn visit_none<E: de::Error>(self) -> Result<Self::Value, E> {
Ok(None)
}
fn visit_unit<E: de::Error>(self) -> Result<Self::Value, E> {
Ok(None)
}
fn visit_some<D: Deserializer<'de>>(self, d: D) -> Result<Self::Value, D::Error> {
super::deserialize_int(d).map(Some)
}
}
d.deserialize_option(OptVisitor(std::marker::PhantomData))
}
}
pub mod vec_int {
use super::*;
pub fn serialize<T: Display, S: Serializer>(v: &[T], s: S) -> Result<S::Ok, S::Error> {
use serde::ser::SerializeSeq;
let mut seq = s.serialize_seq(Some(v.len()))?;
for item in v {
seq.serialize_element(&item.to_string())?;
}
seq.end()
}
pub fn deserialize<'de, T, D>(d: D) -> Result<Vec<T>, D::Error>
where
T: FromStr + TryFrom<i64> + TryFrom<u64>,
<T as FromStr>::Err: Display,
D: Deserializer<'de>,
{
struct SeqVisitor<T>(std::marker::PhantomData<T>);
impl<'de, T> Visitor<'de> for SeqVisitor<T>
where
T: FromStr + TryFrom<i64> + TryFrom<u64>,
<T as FromStr>::Err: Display,
{
type Value = Vec<T>;
fn expecting(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("a sequence of integers")
}
fn visit_unit<E: de::Error>(self) -> Result<Self::Value, E> {
Ok(Vec::new())
}
fn visit_seq<A: de::SeqAccess<'de>>(self, mut a: A) -> Result<Self::Value, A::Error> {
let mut out = Vec::with_capacity(a.size_hint().unwrap_or(0));
while let Some(v) = a.next_element_seed(IntSeed(std::marker::PhantomData))? {
out.push(v);
}
Ok(out)
}
}
struct IntSeed<T>(std::marker::PhantomData<T>);
impl<'de, T> de::DeserializeSeed<'de> for IntSeed<T>
where
T: FromStr + TryFrom<i64> + TryFrom<u64>,
<T as FromStr>::Err: Display,
{
type Value = T;
fn deserialize<D: Deserializer<'de>>(self, d: D) -> Result<T, D::Error> {
super::deserialize_int(d)
}
}
d.deserialize_any(SeqVisitor(std::marker::PhantomData))
}
}
pub mod opt_bytes {
use super::*;
pub fn serialize<S: Serializer>(v: &Option<Vec<u8>>, s: S) -> Result<S::Ok, S::Error> {
match v {
Some(b) => s.serialize_str(&STANDARD.encode(b)),
None => s.serialize_none(),
}
}
pub fn deserialize<'de, D: Deserializer<'de>>(d: D) -> Result<Option<Vec<u8>>, D::Error> {
let raw = Option::<String>::deserialize(d)?;
raw.map(|s| decode_base64(&s)).transpose()
}
}
pub mod vec_bytes {
use super::*;
pub fn serialize<S: Serializer>(v: &[Vec<u8>], s: S) -> Result<S::Ok, S::Error> {
use serde::ser::SerializeSeq;
let mut seq = s.serialize_seq(Some(v.len()))?;
for item in v {
seq.serialize_element(&STANDARD.encode(item))?;
}
seq.end()
}
pub fn deserialize<'de, D: Deserializer<'de>>(d: D) -> Result<Vec<Vec<u8>>, D::Error> {
let raw = Option::<Vec<String>>::deserialize(d)?.unwrap_or_default();
raw.iter().map(|s| decode_base64(s)).collect()
}
}
#[cfg(test)]
mod tests {
use serde::{Deserialize, Serialize};
#[derive(Debug, PartialEq, Serialize, Deserialize)]
struct Sample {
#[serde(
default,
with = "super::opt_int",
skip_serializing_if = "Option::is_none"
)]
seq: Option<i64>,
#[serde(
default,
with = "super::opt_bytes",
skip_serializing_if = "Option::is_none"
)]
data: Option<Vec<u8>>,
#[serde(
default,
with = "super::vec_int",
skip_serializing_if = "Vec::is_empty"
)]
counts: Vec<u64>,
}
#[test]
fn accepts_string_and_numeric_integers() {
let from_string: Sample = serde_json::from_str(r#"{"seq":"17"}"#).unwrap();
let from_number: Sample = serde_json::from_str(r#"{"seq":17}"#).unwrap();
assert_eq!(from_string.seq, Some(17));
assert_eq!(from_number.seq, from_string.seq);
}
#[test]
fn emits_the_canonical_string_form() {
let s = Sample {
seq: Some(-3),
data: None,
counts: vec![1, 2],
};
assert_eq!(
serde_json::to_string(&s).unwrap(),
r#"{"seq":"-3","counts":["1","2"]}"#
);
}
#[test]
fn absent_fields_stay_absent() {
let s: Sample = serde_json::from_str("{}").unwrap();
assert_eq!(
s,
Sample {
seq: None,
data: None,
counts: vec![]
}
);
}
#[test]
fn base64_round_trips_across_alphabets() {
let padded: Sample = serde_json::from_str(r#"{"data":"//79"}"#).unwrap();
let url_safe: Sample = serde_json::from_str(r#"{"data":"__79"}"#).unwrap();
assert_eq!(padded.data, Some(vec![0xff, 0xfe, 0xfd]));
assert_eq!(url_safe.data, padded.data);
assert_eq!(
serde_json::to_string(&padded).unwrap(),
r#"{"data":"//79"}"#
);
}
#[test]
fn null_decodes_as_absent() {
let s: Sample = serde_json::from_str(r#"{"seq":null,"data":null}"#).unwrap();
assert_eq!(s.seq, None);
assert_eq!(s.data, None);
}
}