use std::collections::{BTreeMap, BTreeSet};
use proc_macro2::{Ident, TokenStream};
use quote::{format_ident, quote};
use crate::build::{
JsonValue, StructProperty, Type, TypeCommon, TypeCommonBuilt, check_properties, validate_ident,
};
use crate::default::EnumDefault;
use crate::error::{Error, NameAxis};
use crate::output::Outputspace;
use crate::serde_attrs::SerdeDerives;
use crate::{TypespaceBuilder, TypespaceRenderer, TypespaceTrait, TypespaceTraitSet};
#[derive(Debug, Clone)]
pub struct Enum<Id> {
pub(crate) common: TypeCommon,
pub(crate) tag_type: Option<EnumTagType>,
pub(crate) variants: Vec<EnumVariant<Id>>,
pub(crate) deny_unknown_fields: bool,
}
impl<Id> Default for Enum<Id> {
fn default() -> Self {
Self::new()
}
}
impl<Id> Enum<Id> {
pub fn new() -> Self {
Self {
common: Default::default(),
tag_type: None,
variants: Vec::new(),
deny_unknown_fields: false,
}
}
pub fn name(mut self, name: impl Into<String>) -> Self {
self.common.name = Some(name.into());
self
}
pub fn description(mut self, description: impl Into<String>) -> Self {
self.common.description = Some(description.into());
self
}
pub fn default(mut self, default: impl Into<JsonValue>) -> Self {
self.common.default = Some(default.into());
self
}
pub fn extra_derives(mut self, derives: impl IntoIterator<Item = impl Into<String>>) -> Self {
self.common
.extra_derives
.extend(derives.into_iter().map(Into::into));
self
}
pub fn extra_attrs(mut self, attrs: impl IntoIterator<Item = impl Into<String>>) -> Self {
self.common
.extra_attrs
.extend(attrs.into_iter().map(Into::into));
self
}
pub fn tag_type(mut self, tag_type: EnumTagType) -> Self {
self.tag_type = Some(tag_type);
self
}
pub fn variants(mut self, variants: impl IntoIterator<Item = EnumVariant<Id>>) -> Self {
self.variants.extend(variants);
self
}
pub fn deny_unknown_fields(mut self) -> Self {
self.deny_unknown_fields = true;
self
}
pub fn build(self) -> Result<Type<Id>, Error<Id>>
where
Id: std::fmt::Debug + std::fmt::Display,
{
self.validate()?;
Ok(Type::Enum(self))
}
pub(crate) fn validate(&self) -> Result<(), Error<Id>>
where
Id: std::fmt::Debug + std::fmt::Display,
{
self.common.validate_name("enum")?;
let type_name = self.common.built_name();
if self.tag_type.is_none() {
return Err(Error::MissingTagType {
name: type_name.to_string(),
});
}
let mut rust_names = BTreeSet::new();
let mut wire_names = BTreeSet::new();
for variant in &self.variants {
validate_ident("variant", &variant.rust_name)?;
if !rust_names.insert(variant.rust_name.clone()) {
return Err(Error::DuplicateItemName {
kind: "variant",
type_name: type_name.to_string(),
name: variant.rust_name.clone(),
axis: NameAxis::Rust,
});
}
let wire_name = variant
.rename
.clone()
.unwrap_or_else(|| variant.rust_name.clone());
if !wire_names.insert(wire_name.clone()) {
return Err(Error::DuplicateItemName {
kind: "variant",
type_name: type_name.to_string(),
name: wire_name,
axis: NameAxis::Wire,
});
}
if let VariantDetails::Struct(properties) = &variant.details {
check_properties(&format!("{type_name}::{}", variant.rust_name), properties)?;
}
}
Ok(())
}
pub fn get_name(&self) -> Option<&str> {
self.common.name()
}
pub fn get_description(&self) -> Option<&str> {
self.common.description()
}
pub fn get_default(&self) -> Option<&serde_json::Value> {
self.common.default()
}
pub fn get_extra_derives(&self) -> &[String] {
self.common.extra_derives()
}
pub fn get_extra_attrs(&self) -> &[String] {
self.common.extra_attrs()
}
pub fn get_tag_type(&self) -> Option<&EnumTagType> {
self.tag_type.as_ref()
}
pub fn get_variants(&self) -> &[EnumVariant<Id>] {
&self.variants
}
pub fn get_deny_unknown_fields(&self) -> bool {
self.deny_unknown_fields
}
pub fn all_tagged_unit_variants(&self) -> bool {
self.tag_type
.as_ref()
.is_some_and(|tag_type| *tag_type != EnumTagType::Untagged)
&& !self.variants.is_empty()
&& self
.variants
.iter()
.all(|variant| matches!(variant.details, VariantDetails::Unit))
}
pub fn all_untagged_item_variants(&self) -> bool {
self.tag_type
.as_ref()
.is_some_and(|tag_type| *tag_type == EnumTagType::Untagged)
&& !self.variants.is_empty()
&& self
.variants
.iter()
.all(|variant| matches!(variant.details, VariantDetails::Item(_)))
}
pub(crate) fn every_variant_is_unit(&self) -> bool {
self.variants
.iter()
.all(|variant| matches!(variant.details, VariantDetails::Unit))
}
pub(crate) fn check_field_defaults(
&self,
typespace: &TypespaceBuilder<Id>,
) -> Result<BTreeSet<Id>, Error<Id>>
where
Id: Clone + Ord + std::fmt::Debug + std::fmt::Display,
{
self.variants
.iter()
.try_fold(BTreeSet::new(), |natives, variant| match &variant.details {
VariantDetails::Struct(items) => {
items.iter().try_fold(natives, |mut natives, prop| {
natives.extend(prop.check_defaults(typespace)?);
Ok(natives)
})
}
_ => Ok(natives),
})
}
}
#[derive(Default)]
struct EnumSpecialImpls {
display_impl: TokenStream,
from_str_impl: TokenStream,
try_from_impl: TokenStream,
}
impl<Id: Clone + Ord + std::fmt::Debug + std::fmt::Display> Enum<Id> {
pub(crate) fn children(&self) -> Vec<Id> {
self.variants
.iter()
.flat_map(|variant| variant.children())
.collect()
}
pub(crate) fn render(
&self,
id: &Id,
typespace: &TypespaceRenderer<'_, Id>,
out: &mut Outputspace,
) -> TokenStream {
let Self {
common:
TypeCommon {
name,
description,
default,
built:
Some(TypeCommonBuilt {
traits,
from_string_irrefutable: _,
}),
extra_derives,
extra_attrs,
},
tag_type,
variants,
deny_unknown_fields,
} = self
else {
unreachable!()
};
let name = name.as_deref().expect("validated type has a name");
let tag_type = tag_type.as_ref().expect("validated enum has a tag type");
let description = description.as_ref().map(|desc| quote! { #[doc = #desc] });
let name_ident = format_ident!("{name}");
let mut derived_traits = traits.clone();
let all_unit_variants = self.all_tagged_unit_variants();
let all_item_variants = self.all_untagged_item_variants();
let special_impls = match (all_unit_variants, all_item_variants) {
(true, true) => unreachable!(),
(true, false) => self.render_tagged_unit_variant_impls(
typespace,
out,
&name_ident,
&mut derived_traits,
),
(false, true) => self.render_untagged_item_variant_impls(
typespace,
out,
&name_ident,
&mut derived_traits,
),
(false, false) => EnumSpecialImpls::default(),
};
let every_variant_is_unit = self.every_variant_is_unit();
let serde_derives = SerdeDerives::new(&derived_traits);
let mut serde = serde_derives.attrs();
serde.extend(match tag_type {
EnumTagType::External => Vec::new(),
EnumTagType::Internal { tag } => vec![quote! { tag = #tag }],
EnumTagType::Adjacent { tag, content } => {
vec![quote! { tag = #tag }, quote! { content = #content }]
}
EnumTagType::Untagged => vec![quote! { untagged }],
});
let variant_from = self.render_variant_from(typespace, &name_ident);
let (default_impl, unit_default_value) =
if derived_traits.contains(&TypespaceTrait::Default) {
let default_value = &default
.as_ref()
.expect("validated type with Default among its impls must have a valid default")
.0;
let generated_default = typespace.generate_default_enum(default_value, id);
match generated_default {
EnumDefault::Value(default_value) => {
let default_impl = quote! {
impl ::std::default::Default for #name_ident {
fn default() -> Self {
#default_value
}
}
};
derived_traits.remove(TypespaceTrait::Default);
(default_impl, None)
}
EnumDefault::Variant(variant_name) => (TokenStream::new(), Some(variant_name)),
}
} else {
(TokenStream::new(), None)
};
let rendered_variants = variants.iter().map(|variant| {
let EnumVariant {
rust_name,
rename,
description,
details,
} = variant;
let variant_ident = format_ident!("{}", rust_name);
let mut variant_serde = serde_derives.attrs();
variant_serde.extend(rename.as_ref().map(|n| quote! { rename = #n }));
let description = description.as_ref().map(|desc| quote! { #[doc = #desc] });
let default_attr = (unit_default_value.as_ref() == Some(rust_name)).then(|| {
quote! { #[default] }
});
let data = match details {
VariantDetails::Unit => TokenStream::new(),
VariantDetails::Item(item) => {
let item_ident = typespace.render_ident(item);
quote! { (#item_ident) }
}
VariantDetails::Tuple(items) => {
let item_idents = items.iter().map(|item| typespace.render_ident(item));
quote! { ( #( #item_idents, )* ) }
}
VariantDetails::Struct(properties) => {
let properties = properties.iter().map(|prop| {
typespace.render_struct_property(
prop,
serde_derives,
false,
&format!("{name}{rust_name}"),
out,
)
});
quote! { { #( #properties, )* } }
}
};
quote! {
#description
#variant_serde
#default_attr
#variant_ident #data
}
});
if serde_derives.deserialize() && *deny_unknown_fields {
serde.push(quote! { deny_unknown_fields });
}
let derive_attr =
typespace.render_derives(&derived_traits, extra_derives, every_variant_is_unit);
let attrs = typespace.render_attrs(extra_attrs);
let EnumSpecialImpls {
display_impl,
from_str_impl,
try_from_impl,
} = special_impls;
quote! {
#description
#( #attrs )*
#derive_attr
#serde
pub enum #name_ident {
#( #rendered_variants, )*
}
#display_impl
#from_str_impl
#try_from_impl
#default_impl
#( #variant_from )*
}
}
fn render_tagged_unit_variant_impls(
&self,
typespace: &TypespaceRenderer<'_, Id>,
out: &mut Outputspace,
name_ident: &Ident,
derived_traits: &mut TypespaceTraitSet,
) -> EnumSpecialImpls {
assert!(
self.all_tagged_unit_variants(),
"{} is not an all-unit-variant enum",
self.common.built_name(),
);
let (variant_idents, variant_names): (Vec<_>, Vec<_>) = self
.variants
.iter()
.map(|variant| (format_ident!("{}", variant.rust_name), variant.json_name()))
.unzip();
let display_impl = if derived_traits.remove(TypespaceTrait::Display) {
quote! {
impl ::std::fmt::Display for #name_ident {
fn fmt(&self, f: &mut ::std::fmt::Formatter<'_>)
-> ::std::fmt::Result
{
match *self {
#( Self::#variant_idents => f.write_str(#variant_names), )*
}
}
}
}
} else {
Default::default()
};
let (from_str_impl, try_from_impl) = if derived_traits.remove(TypespaceTrait::FromStr) {
typespace.add_error_mod(out);
let string_type = typespace.render_std_string();
let from_str_impl = quote! {
impl ::std::str::FromStr for #name_ident {
type Err = self::error::ConversionError;
fn from_str(value: &str)
-> ::std::result::Result<Self, self::error::ConversionError>
{
match value {
#( #variant_names => Ok(Self::#variant_idents), )*
_ => Err("invalid value".into()),
}
}
}
};
let try_from_impl = quote! {
impl ::std::convert::TryFrom<&str> for #name_ident {
type Error = self::error::ConversionError;
fn try_from(value: &str)
-> ::std::result::Result<Self, self::error::ConversionError>
{
value.parse()
}
}
impl ::std::convert::TryFrom<#string_type> for #name_ident {
type Error = self::error::ConversionError;
fn try_from(value: #string_type)
-> ::std::result::Result<Self, self::error::ConversionError>
{
value.parse()
}
}
};
(from_str_impl, try_from_impl)
} else {
(TokenStream::new(), TokenStream::new())
};
EnumSpecialImpls {
display_impl,
from_str_impl,
try_from_impl,
}
}
fn render_untagged_item_variant_impls(
&self,
typespace: &TypespaceRenderer<'_, Id>,
out: &mut Outputspace,
name_ident: &Ident,
derived_traits: &mut TypespaceTraitSet,
) -> EnumSpecialImpls {
let variant_idents = self
.variants
.iter()
.map(|variant| format_ident!("{}", variant.rust_name))
.collect::<Vec<_>>();
let (from_str_impl, try_from_impl) = if derived_traits.remove(TypespaceTrait::FromStr) {
typespace.add_error_mod(out);
let string_type = typespace.render_std_string();
let from_str_impl = quote! {
impl ::std::str::FromStr for #name_ident {
type Err = self::error::ConversionError;
fn from_str(value: &str) ->
::std::result::Result<Self, self::error::ConversionError>
{
#(
if let Ok(v) = value.parse() {
Ok(Self::#variant_idents(v))
} else
)*
{
Err("string conversion failed for all variants".into())
}
}
}
};
let try_from_impl = quote! {
impl ::std::convert::TryFrom<&str> for #name_ident {
type Error = self::error::ConversionError;
fn try_from(value: &str) ->
::std::result::Result<Self, self::error::ConversionError>
{
value.parse()
}
}
impl ::std::convert::TryFrom<#string_type> for #name_ident {
type Error = self::error::ConversionError;
fn try_from(value: #string_type)
-> ::std::result::Result<Self, self::error::ConversionError>
{
value.parse()
}
}
};
(from_str_impl, try_from_impl)
} else {
(TokenStream::new(), TokenStream::new())
};
let display_impl = if derived_traits.remove(TypespaceTrait::Display) {
quote! {
impl ::std::fmt::Display for #name_ident {
fn fmt(&self, f: &mut ::std::fmt::Formatter<'_>) -> ::std::fmt::Result {
match self {
#(Self::#variant_idents(x) => x.fmt(f),)*
}
}
}
}
} else {
Default::default()
};
EnumSpecialImpls {
display_impl,
from_str_impl,
try_from_impl,
}
}
fn render_variant_from(
&self,
typespace: &TypespaceRenderer<'_, Id>,
name_ident: &Ident,
) -> Vec<TokenStream> {
let unique_variants =
self.variants
.iter()
.enumerate()
.fold(BTreeMap::new(), |mut map, (index, variant)| {
let key = match &variant.details {
VariantDetails::Item(id) => {
vec![typespace.render_ident(id).to_string()]
}
VariantDetails::Tuple(ids) => ids
.iter()
.map(|id| typespace.render_ident(id).to_string())
.collect::<Vec<_>>(),
VariantDetails::Unit | VariantDetails::Struct(_) => return map,
};
map.entry(key)
.and_modify(|seen| *seen = None)
.or_insert(Some((index, variant)));
map
});
unique_variants
.into_values()
.flatten()
.collect::<BTreeMap<_, _>>()
.into_values()
.filter_map(|variant| {
let variant_ident = format_ident!("{}", variant.rust_name);
match &variant.details {
VariantDetails::Item(id)
if matches!(typespace.types.get(id), Some(Type::String)) =>
{
None
}
VariantDetails::Item(id) => {
let payload = typespace.render_ident(id);
Some(quote! {
impl ::std::convert::From<#payload> for #name_ident {
fn from(value: #payload) -> Self {
Self::#variant_ident(value)
}
}
})
}
VariantDetails::Tuple(ids) => {
let payloads = ids
.iter()
.map(|id| typespace.render_ident(id))
.collect::<Vec<_>>();
let payload = match ids.len() {
1 => quote! { ( #( #payloads, )* ) },
_ => quote! { ( #( #payloads ),* ) },
};
let field = (0..ids.len()).map(syn::Index::from);
Some(quote! {
impl ::std::convert::From<#payload> for #name_ident {
fn from(value: #payload) -> Self {
Self::#variant_ident( #( value.#field, )* )
}
}
})
}
VariantDetails::Unit | VariantDetails::Struct(_) => None,
}
})
.collect::<Vec<_>>()
}
}
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord)]
pub enum EnumTagType {
External,
Internal { tag: String },
Adjacent { tag: String, content: String },
Untagged,
}
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord)]
#[non_exhaustive]
pub struct EnumVariant<Id> {
pub(crate) rust_name: String,
pub(crate) rename: Option<String>,
pub(crate) description: Option<String>,
pub(crate) details: VariantDetails<Id>,
}
impl<Id> EnumVariant<Id> {
pub fn new(rust_name: impl Into<String>, details: VariantDetails<Id>) -> Self {
Self {
rust_name: rust_name.into(),
rename: None,
description: None,
details,
}
}
pub fn with_rename(mut self, rename: impl Into<String>) -> Self {
self.rename = Some(rename.into());
self
}
pub fn with_description(mut self, description: impl Into<String>) -> Self {
self.description = Some(description.into());
self
}
pub fn rust_name(&self) -> &str {
&self.rust_name
}
pub fn rename(&self) -> Option<&str> {
self.rename.as_deref()
}
pub fn json_name(&self) -> &str {
self.rename.as_deref().unwrap_or(&self.rust_name)
}
pub fn description(&self) -> Option<&str> {
self.description.as_deref()
}
pub fn details(&self) -> &VariantDetails<Id> {
&self.details
}
pub(crate) fn relation(&self) -> crate::error::Relation {
crate::error::Relation::Variant(self.rust_name.clone())
}
}
impl<Id: Clone> EnumVariant<Id> {
fn children(&self) -> Vec<Id> {
match &self.details {
VariantDetails::Unit => Vec::new(),
VariantDetails::Item(id) => vec![id.clone()],
VariantDetails::Tuple(items) => items.clone(),
VariantDetails::Struct(items) => {
items.iter().map(|prop| prop.type_id.clone()).collect()
}
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord)]
pub enum VariantDetails<Id> {
Unit,
Item(Id),
Tuple(Vec<Id>),
Struct(Vec<StructProperty<Id>>),
}