use std::fmt;
use std::marker::PhantomData;
use serde::de::{self, Deserializer, Visitor};
use serde::{Deserialize, Serialize, Serializer};
pub(super) trait Tagged {
const TAG: &'static str;
}
pub(super) struct Tag<T>(PhantomData<fn() -> T>);
impl<T> Default for Tag<T> {
fn default() -> Self {
Tag(PhantomData)
}
}
impl<T> Clone for Tag<T> {
fn clone(&self) -> Self {
*self
}
}
impl<T> Copy for Tag<T> {}
impl<T> PartialEq for Tag<T> {
fn eq(&self, _: &Self) -> bool {
true
}
}
impl<T: Tagged> fmt::Debug for Tag<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(T::TAG)
}
}
impl<T: Tagged> Serialize for Tag<T> {
fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
serializer.serialize_str(T::TAG)
}
}
impl<'de, T: Tagged> Deserialize<'de> for Tag<T> {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
struct TagVisitor<T>(PhantomData<fn() -> T>);
impl<T: Tagged> Visitor<'_> for TagVisitor<T> {
type Value = Tag<T>;
fn expecting(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "\"{}\"", T::TAG)
}
fn visit_str<E: de::Error>(self, v: &str) -> Result<Self::Value, E> {
if v == T::TAG {
Ok(Tag::default())
} else {
Err(E::invalid_value(de::Unexpected::Str(v), &self))
}
}
}
deserializer.deserialize_str(TagVisitor(PhantomData))
}
}
macro_rules! tagged_enum {
(
$(#[$meta:meta])*
$vis:vis enum $name:ident($key:literal) {
$($variant:ident($ty:ty),)*
}
) => {
$(#[$meta])*
#[derive(Debug, Clone, PartialEq)]
$vis enum $name {
$($variant($ty),)*
Unknown(::serde_json::Value),
}
$(
impl From<$ty> for $name {
fn from(value: $ty) -> Self {
$name::$variant(value)
}
}
)*
impl ::serde::Serialize for $name {
fn serialize<S: ::serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
match self {
$($name::$variant(value) => value.serialize(serializer),)*
$name::Unknown(value) => value.serialize(serializer),
}
}
}
impl<'de> ::serde::Deserialize<'de> for $name {
fn deserialize<D: ::serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let value = ::serde_json::Value::deserialize(deserializer)?;
let Some(tag) = value.get($key).and_then(::serde_json::Value::as_str) else {
return Ok($name::Unknown(value));
};
$(
if tag == <$ty as super::tag::Tagged>::TAG {
return ::serde_json::from_value(value)
.map($name::$variant)
.map_err(::serde::de::Error::custom);
}
)*
Ok($name::Unknown(value))
}
}
};
}
pub(super) use tagged_enum;