#![deny(missing_docs)]
use std::collections::HashSet;
use proc_macro::TokenStream;
use proc_macro2::{Span, TokenStream as TokenStream2};
use quote::{format_ident, quote};
use syn::ext::IdentExt;
use syn::spanned::Spanned;
use syn::visit::{self, Visit};
use syn::{
Attribute, Data, DataEnum, DeriveInput, Error, Fields, FnArg, GenericArgument, Ident,
ItemTrait, Lit, Meta, Pat, PathArguments, ReturnType, TraitItem, TraitItemFn, Type,
parse_macro_input,
};
#[proc_macro_attribute]
pub fn interface(attribute: TokenStream, item: TokenStream) -> TokenStream {
let attribute = TokenStream2::from(attribute);
let item = parse_macro_input!(item as ItemTrait);
match expand_interface(attribute, item) {
Ok(tokens) => tokens.into(),
Err(error) => error.into_compile_error().into(),
}
}
#[proc_macro_derive(WipType)]
pub fn derive_wip_type(item: TokenStream) -> TokenStream {
let item = parse_macro_input!(item as DeriveInput);
match expand_wip_type(item) {
Ok(tokens) => tokens.into(),
Err(error) => error.into_compile_error().into(),
}
}
#[derive(Clone)]
struct Parameter {
ident: Ident,
name: String,
ty: Type,
optional: bool,
documentation: Option<Documentation>,
}
struct Operation {
ident: Ident,
name: String,
arguments_ident: Ident,
parameters: Vec<Parameter>,
result_ty: Type,
unit_result: bool,
documentation: Option<Documentation>,
return_documentation: Option<Documentation>,
}
#[derive(Clone)]
struct Documentation {
summary: String,
details: Option<String>,
}
fn expand_interface(attribute: TokenStream2, item: ItemTrait) -> syn::Result<TokenStream2> {
if !attribute.is_empty() {
return Err(Error::new_spanned(
attribute,
"the interface attribute does not accept arguments",
));
}
validate_interface_header(&item)?;
let trait_ident = &item.ident;
let definition_ident = format_ident!("{}Definition", trait_ident, span = trait_ident.span());
let trait_documentation = documentation(&item.attrs);
let mut operation_names = HashSet::new();
let mut operations = Vec::new();
for trait_item in &item.items {
let method = match trait_item {
TraitItem::Fn(method) => method,
TraitItem::Type(item) => {
return Err(Error::new_spanned(
item,
"associated types are not supported in WIP interfaces",
));
}
TraitItem::Const(item) => {
return Err(Error::new_spanned(
item,
"associated constants are not supported in WIP interfaces",
));
}
other => {
return Err(Error::new_spanned(
other,
"only operation methods are supported in WIP interfaces",
));
}
};
let operation = parse_operation(trait_ident, method)?;
if !operation_names.insert(operation.name.clone()) {
return Err(Error::new(
operation.ident.span(),
format!("duplicate operation name `{}`", operation.name),
));
}
operations.push(operation);
}
let mut seen_argument_items = HashSet::new();
let mut colliding_argument_items = HashSet::new();
for operation in &operations {
let name = operation.arguments_ident.to_string();
if !seen_argument_items.insert(name.clone()) {
colliding_argument_items.insert(name);
}
}
for operation in &mut operations {
if colliding_argument_items.contains(&operation.arguments_ident.to_string()) {
let suffix = operation
.name
.as_bytes()
.iter()
.map(|byte| format!("{byte:02x}"))
.collect::<String>();
operation.arguments_ident = format_ident!(
"{}{}ArgumentsN{}",
trait_ident,
pascal_case(&operation.name),
suffix,
span = operation.ident.span()
);
}
}
let argument_items = operations.iter().map(argument_item);
let operation_declarations = operations.iter().map(operation_declaration);
let declaration_collectors = operations.iter().flat_map(operation_collectors);
let dispatch_arms = operations.iter().map(dispatch_arm);
let trait_doc = documentation_tokens(trait_documentation.as_ref());
let mut emitted_item = item.clone();
strip_parameter_documentation(&mut emitted_item);
Ok(quote! {
#emitted_item
#[doc = concat!("Generated WIP interface definition for [`", stringify!(#trait_ident), "`].")]
pub struct #definition_ident;
impl ::wip_interface::InterfaceDefinition for #definition_ident {
fn descriptor() -> ::wip_interface::wip_protocol::InterfaceDescriptor {
let mut declarations = ::std::vec::Vec::new();
#(#declaration_collectors)*
::wip_interface::wip_protocol::InterfaceDescriptor {
format: ::std::string::String::from(
::wip_interface::wip_protocol::INTERFACE_FORMAT_V1,
),
documentation: #trait_doc,
types: declarations,
operations: ::std::vec![#(#operation_declarations),*],
}
}
}
impl<T: #trait_ident> ::wip_interface::InterfaceImplementation<T> for #definition_ident {
fn dispatch(
implementation: &T,
operation: &str,
arguments: ::std::collections::BTreeMap<
::std::string::String,
::wip_interface::wip_protocol::Value,
>,
) -> ::std::result::Result<
::wip_interface::wip_protocol::Value,
::wip_interface::DispatchError,
> {
let descriptor = Self::descriptor();
match operation {
#(#dispatch_arms)*
_ => ::std::result::Result::Err(
::wip_interface::DispatchError::UnknownOperation(
::std::string::String::from(operation),
),
),
}
}
}
impl #definition_ident {
pub fn descriptor() -> ::wip_interface::wip_protocol::InterfaceDescriptor {
<Self as ::wip_interface::InterfaceDefinition>::descriptor()
}
pub fn dispatch<T: #trait_ident>(
implementation: &T,
operation: &str,
arguments: ::std::collections::BTreeMap<
::std::string::String,
::wip_interface::wip_protocol::Value,
>,
) -> ::std::result::Result<
::wip_interface::wip_protocol::Value,
::wip_interface::DispatchError,
> {
<Self as ::wip_interface::InterfaceImplementation<T>>::dispatch(
implementation,
operation,
arguments,
)
}
}
#(#argument_items)*
})
}
fn strip_parameter_documentation(item: &mut ItemTrait) {
for trait_item in &mut item.items {
let TraitItem::Fn(method) = trait_item else {
continue;
};
for input in &mut method.sig.inputs {
if let FnArg::Typed(typed) = input {
typed
.attrs
.retain(|attribute| !attribute.path().is_ident("doc"));
}
}
}
}
fn validate_interface_header(item: &ItemTrait) -> syn::Result<()> {
if !matches!(item.vis, syn::Visibility::Public(_)) {
return Err(Error::new_spanned(
&item.vis,
"a WIP interface trait must be public",
));
}
if item.unsafety.is_some() {
return Err(Error::new_spanned(
item.unsafety,
"unsafe interface traits are not supported",
));
}
if item.auto_token.is_some() {
return Err(Error::new_spanned(
item.auto_token,
"auto traits are not supported as WIP interfaces",
));
}
if !item.generics.params.is_empty() || item.generics.where_clause.is_some() {
return Err(Error::new_spanned(
&item.generics,
"generic WIP interface traits are not supported",
));
}
if item.colon_token.is_some() || !item.supertraits.is_empty() {
return Err(Error::new_spanned(
&item.supertraits,
"WIP interface trait inheritance is not supported",
));
}
Ok(())
}
fn parse_operation(trait_ident: &Ident, method: &TraitItemFn) -> syn::Result<Operation> {
let signature = &method.sig;
if signature.constness.is_some() {
return Err(Error::new_spanned(
signature.constness,
"const operations are not supported",
));
}
if signature.asyncness.is_some() {
return Err(Error::new_spanned(
signature.asyncness,
"async operations are not supported",
));
}
if signature.unsafety.is_some() {
return Err(Error::new_spanned(
signature.unsafety,
"unsafe operations are not supported",
));
}
if signature.abi.is_some() {
return Err(Error::new_spanned(
&signature.abi,
"extern operations are not supported",
));
}
if signature.variadic.is_some() {
return Err(Error::new_spanned(
&signature.variadic,
"variadic operations are not supported",
));
}
if !signature.generics.params.is_empty() || signature.generics.where_clause.is_some() {
return Err(Error::new_spanned(
&signature.generics,
"generic operations are not supported",
));
}
if let Some(default) = &method.default {
return Err(Error::new_spanned(
default,
"WIP operation methods must not have a body",
));
}
let mut inputs = signature.inputs.iter();
let receiver = inputs.next().ok_or_else(|| {
Error::new(
signature.ident.span(),
"a WIP operation must have an `&self` receiver",
)
})?;
validate_receiver(receiver)?;
let mut parameters = Vec::new();
let mut parameter_names = HashSet::new();
for input in inputs {
let typed = match input {
FnArg::Typed(typed) => typed,
FnArg::Receiver(receiver) => {
return Err(Error::new_spanned(
receiver,
"the `&self` receiver must be the first operation argument",
));
}
};
let ident = match typed.pat.as_ref() {
Pat::Ident(pattern)
if pattern.by_ref.is_none()
&& pattern.mutability.is_none()
&& pattern.subpat.is_none() =>
{
pattern.ident.clone()
}
pattern => {
return Err(Error::new_spanned(
pattern,
"operation parameters must be simple named parameters",
));
}
};
let name = source_name(&ident);
if !parameter_names.insert(name.clone()) {
return Err(Error::new(
ident.span(),
format!("duplicate operation parameter `{name}`"),
));
}
let (optional, ty) = optional_type(typed.ty.as_ref(), "operation parameter")?;
validate_value_type(&ty, "operation parameter")?;
parameters.push(Parameter {
ident,
name,
ty,
optional,
documentation: documentation(&typed.attrs),
});
}
let result_ty = parse_operation_result(&signature.output)?;
let unit_result = is_unit(&result_ty);
if !unit_result {
validate_value_type(&result_ty, "operation return type")?;
}
let name = source_name(&signature.ident);
let method_docs = doc_lines(&method.attrs);
let (operation_doc_lines, return_doc_lines) = split_returns_section(method_docs);
let pascal_name = pascal_case(&name);
let arguments_ident = format_ident!(
"{}{}Arguments",
trait_ident,
pascal_name,
span = signature.ident.span()
);
Ok(Operation {
ident: signature.ident.clone(),
name,
arguments_ident,
parameters,
result_ty,
unit_result,
documentation: documentation_from_lines(operation_doc_lines),
return_documentation: documentation_from_lines(return_doc_lines),
})
}
fn validate_receiver(input: &FnArg) -> syn::Result<()> {
let receiver = match input {
FnArg::Receiver(receiver) => receiver,
FnArg::Typed(typed) => {
return Err(Error::new_spanned(
typed,
"a WIP operation must start with an `&self` receiver",
));
}
};
let reference = receiver.reference.as_ref();
let is_plain_shared_reference = reference.is_some()
&& receiver.mutability.is_none()
&& receiver.colon_token.is_none()
&& reference
.and_then(|(_, lifetime)| lifetime.as_ref())
.is_none();
if !is_plain_shared_reference {
return Err(Error::new_spanned(
receiver,
"the only supported receiver is exactly `&self`",
));
}
Ok(())
}
fn parse_operation_result(output: &ReturnType) -> syn::Result<Type> {
let ty = match output {
ReturnType::Default => {
return Err(Error::new_spanned(
output,
"an operation must explicitly return `OperationResult<T>`",
));
}
ReturnType::Type(_, ty) => ty.as_ref(),
};
reject_custom_lifetimes(ty, "operation return type")?;
let path = match ty {
Type::Path(path) if path.qself.is_none() => &path.path,
_ => {
return Err(Error::new_spanned(
ty,
"an operation must explicitly return `OperationResult<T>`",
));
}
};
let segment = path.segments.last().ok_or_else(|| {
Error::new_spanned(
ty,
"an operation must explicitly return `OperationResult<T>`",
)
})?;
if segment.ident != "OperationResult" {
return Err(Error::new_spanned(
ty,
"an operation must explicitly return `OperationResult<T>`",
));
}
let arguments = match &segment.arguments {
PathArguments::AngleBracketed(arguments) => arguments,
_ => {
return Err(Error::new_spanned(
segment,
"`OperationResult` must have exactly one application result type",
));
}
};
let mut types = arguments.args.iter().filter_map(|argument| match argument {
GenericArgument::Type(ty) => Some(ty),
_ => None,
});
let result = types.next().cloned();
if arguments.args.len() != 1 || types.next().is_some() {
return Err(Error::new_spanned(
arguments,
"`OperationResult` must have exactly one application result type",
));
}
let result = result.ok_or_else(|| {
Error::new_spanned(
arguments,
"`OperationResult` must have exactly one application result type",
)
})?;
if direct_option(&result)?.is_some() || contains_option(&result) {
return Err(Error::new_spanned(
result,
"`Option<T>` is only supported directly on DTO fields and operation parameters",
));
}
Ok(result)
}
fn argument_item(operation: &Operation) -> TokenStream2 {
let arguments_ident = &operation.arguments_ident;
let fields = operation.parameters.iter().map(|parameter| {
let ident = ¶meter.ident;
let name = ¶meter.name;
let ty = parameter_rust_type(parameter);
let documentation = format!("Decoded `{name}` operation argument.");
quote!(
#[doc = #documentation]
pub #ident: #ty
)
});
let decoders = operation.parameters.iter().map(|parameter| {
let ident = ¶meter.ident;
let name = ¶meter.name;
let ty = ¶meter.ty;
if parameter.optional {
quote!(#ident: decoder.optional::<#ty>(#name)?)
} else {
quote!(#ident: decoder.required::<#ty>(#name)?)
}
});
let result_ty = &operation.result_ty;
let encode_result = if operation.unit_result {
quote! {
let () = result;
::std::result::Result::Ok(::wip_interface::wip_protocol::Value::Unit)
}
} else {
quote! {
<#result_ty as ::wip_interface::WipType>::encode(result)
}
};
quote! {
#[doc = concat!("Decoded arguments for `", stringify!(#arguments_ident), "`'s operation.")]
pub struct #arguments_ident {
#(#fields,)*
}
impl #arguments_ident {
pub fn decode(
arguments: ::std::collections::BTreeMap<
::std::string::String,
::wip_interface::wip_protocol::Value,
>,
) -> ::std::result::Result<Self, ::wip_interface::CodecError> {
let mut decoder = ::wip_interface::__private::RecordDecoder::new(arguments);
let decoded = Self {
#(#decoders,)*
};
decoder.finish()?;
::std::result::Result::Ok(decoded)
}
pub fn encode_result(
result: #result_ty,
) -> ::std::result::Result<
::wip_interface::wip_protocol::Value,
::wip_interface::CodecError,
> {
#encode_result
}
}
}
}
fn operation_declaration(operation: &Operation) -> TokenStream2 {
let name = &operation.name;
let documentation = documentation_tokens(operation.documentation.as_ref());
let parameters = operation.parameters.iter().map(|parameter| {
let parameter_name = ¶meter.name;
let parameter_doc = documentation_tokens(parameter.documentation.as_ref());
let ty = ¶meter.ty;
let optional = parameter.optional;
quote! {
::wip_interface::wip_protocol::ParameterDeclaration {
name: ::std::string::String::from(#parameter_name),
required: !#optional,
documentation: #parameter_doc,
r#type: <#ty as ::wip_interface::WipType>::type_expr(),
}
}
});
let result_ty = &operation.result_ty;
let return_doc = documentation_tokens(operation.return_documentation.as_ref());
let returns = quote! {
::wip_interface::wip_protocol::ReturnDeclaration {
documentation: #return_doc,
r#type: <#result_ty as ::wip_interface::WipType>::type_expr(),
}
};
quote! {
::wip_interface::wip_protocol::OperationDeclaration {
name: ::std::string::String::from(#name),
documentation: #documentation,
parameters: ::std::vec![#(#parameters),*],
returns: #returns,
}
}
}
fn operation_collectors(operation: &Operation) -> Vec<TokenStream2> {
let mut collectors = operation
.parameters
.iter()
.map(|parameter| {
let ty = ¶meter.ty;
quote!(<#ty as ::wip_interface::WipType>::collect_declarations(&mut declarations);)
})
.collect::<Vec<_>>();
if !operation.unit_result {
let ty = &operation.result_ty;
collectors.push(
quote!(<#ty as ::wip_interface::WipType>::collect_declarations(&mut declarations);),
);
}
collectors
}
fn dispatch_arm(operation: &Operation) -> TokenStream2 {
let operation_name = &operation.name;
let method_ident = &operation.ident;
let arguments_ident = &operation.arguments_ident;
let arguments = operation.parameters.iter().map(|parameter| {
let ident = ¶meter.ident;
quote!(decoded.#ident)
});
quote! {
#operation_name => {
::wip_interface::__private::validate_arguments(
&descriptor,
#operation_name,
&arguments,
)
.map_err(::wip_interface::DispatchError::InvalidArguments)?;
let decoded = #arguments_ident::decode(arguments)
.map_err(::wip_interface::DispatchError::Decode)?;
::std::panic::catch_unwind(::std::panic::AssertUnwindSafe(|| {
let result = implementation.#method_ident(#(#arguments),*)
.map_err(::wip_interface::DispatchError::Host)?;
let encoded = #arguments_ident::encode_result(result)
.map_err(::wip_interface::DispatchError::Encode)?;
descriptor
.validate_result(#operation_name, &encoded)
.map_err(::wip_interface::DispatchError::InvalidResult)?;
::std::result::Result::Ok(encoded)
}))
.map_err(|_| ::wip_interface::DispatchError::Panic)?
}
}
}
fn parameter_rust_type(parameter: &Parameter) -> TokenStream2 {
let ty = ¶meter.ty;
if parameter.optional {
quote!(::std::option::Option<#ty>)
} else {
quote!(#ty)
}
}
fn expand_wip_type(item: DeriveInput) -> syn::Result<TokenStream2> {
if !item.generics.params.is_empty() || item.generics.where_clause.is_some() {
return Err(Error::new_spanned(
&item.generics,
"generic DTOs are not supported by `WipType`",
));
}
let ident = &item.ident;
let type_name = source_name(ident);
let type_documentation = documentation(&item.attrs);
let implementation = match &item.data {
Data::Struct(data) => {
let fields = match &data.fields {
Fields::Named(fields) => &fields.named,
Fields::Unnamed(fields) => {
return Err(Error::new_spanned(
fields,
"`WipType` supports only structs with named fields",
));
}
Fields::Unit => {
return Err(Error::new_spanned(
&item.ident,
"unit structs are not supported by `WipType`",
));
}
};
derive_record(ident, &type_name, type_documentation.as_ref(), fields)?
}
Data::Enum(data) => derive_enum(ident, &type_name, type_documentation.as_ref(), data)?,
Data::Union(data) => {
return Err(Error::new_spanned(
data.union_token,
"Rust unions are not supported by `WipType`",
));
}
};
Ok(implementation)
}
fn derive_record(
ident: &Ident,
type_name: &str,
type_documentation: Option<&Documentation>,
fields: &syn::punctuated::Punctuated<syn::Field, syn::token::Comma>,
) -> syn::Result<TokenStream2> {
let mut parsed_fields = Vec::new();
let mut names = HashSet::new();
for field in fields {
let Some(field_ident) = field.ident.clone() else {
return Err(Error::new_spanned(
field,
"`WipType` record fields must be named",
));
};
let name = source_name(&field_ident);
if !names.insert(name.clone()) {
return Err(Error::new(
field_ident.span(),
format!("duplicate record field `{name}`"),
));
}
let (optional, ty) = optional_type(&field.ty, "DTO field")?;
validate_value_type(&ty, "DTO field")?;
parsed_fields.push(Parameter {
ident: field_ident,
name,
ty,
optional,
documentation: documentation(&field.attrs),
});
}
let declaration_doc = documentation_tokens(type_documentation);
let field_declarations = parsed_fields.iter().map(|field| {
let name = &field.name;
let documentation = documentation_tokens(field.documentation.as_ref());
let ty = &field.ty;
let optional = field.optional;
quote! {
::wip_interface::wip_protocol::FieldDeclaration {
name: ::std::string::String::from(#name),
required: !#optional,
documentation: #documentation,
r#type: <#ty as ::wip_interface::WipType>::type_expr(),
}
}
});
let dependency_collectors = parsed_fields.iter().map(|field| {
let ty = &field.ty;
quote!(<#ty as ::wip_interface::WipType>::collect_declarations(declarations);)
});
let destructured_fields = parsed_fields.iter().map(|field| &field.ident);
let encoders = parsed_fields.iter().map(|field| {
let field_ident = &field.ident;
let field_name = &field.name;
let ty = &field.ty;
if field.optional {
quote! {
if let ::std::option::Option::Some(value) = #field_ident {
fields.insert(
::std::string::String::from(#field_name),
<#ty as ::wip_interface::WipType>::encode(value)?,
);
}
}
} else {
quote! {
fields.insert(
::std::string::String::from(#field_name),
<#ty as ::wip_interface::WipType>::encode(#field_ident)?,
);
}
}
});
let decoders = parsed_fields.iter().map(|field| {
let field_ident = &field.ident;
let field_name = &field.name;
let ty = &field.ty;
if field.optional {
quote!(#field_ident: decoder.optional::<#ty>(#field_name)?)
} else {
quote!(#field_ident: decoder.required::<#ty>(#field_name)?)
}
});
Ok(quote! {
impl ::wip_interface::WipType for #ident {
fn type_expr() -> ::wip_interface::wip_protocol::TypeExpr {
::wip_interface::wip_protocol::TypeExpr::Named {
name: ::std::string::String::from(#type_name),
}
}
fn collect_declarations(
declarations: &mut ::std::vec::Vec<
::wip_interface::wip_protocol::TypeDeclaration,
>,
) {
let declaration = ::wip_interface::wip_protocol::TypeDeclaration {
name: ::std::string::String::from(#type_name),
documentation: #declaration_doc,
definition: ::wip_interface::wip_protocol::TypeExpr::Record {
fields: ::std::vec![#(#field_declarations),*],
},
};
if ::wip_interface::__private::register_declaration(declarations, declaration) {
#(#dependency_collectors)*
}
}
fn encode(
self,
) -> ::std::result::Result<
::wip_interface::wip_protocol::Value,
::wip_interface::CodecError,
> {
let Self { #(#destructured_fields),* } = self;
let mut fields = ::std::collections::BTreeMap::new();
#(#encoders)*
::std::result::Result::Ok(
::wip_interface::wip_protocol::Value::Record(fields),
)
}
fn decode(
value: ::wip_interface::wip_protocol::Value,
) -> ::std::result::Result<Self, ::wip_interface::CodecError> {
match value {
::wip_interface::wip_protocol::Value::Record(fields) => {
let mut decoder = ::wip_interface::__private::RecordDecoder::new(fields);
let decoded = Self {
#(#decoders,)*
};
decoder.finish()?;
::std::result::Result::Ok(decoded)
}
actual => ::std::result::Result::Err(
::wip_interface::__private::decode_type_mismatch("record", actual),
),
}
}
}
})
}
fn derive_enum(
ident: &Ident,
type_name: &str,
type_documentation: Option<&Documentation>,
data: &DataEnum,
) -> syn::Result<TokenStream2> {
for variant in &data.variants {
if let Some((_, discriminant)) = &variant.discriminant {
return Err(Error::new_spanned(
discriminant,
"explicit enum discriminants are not supported by `WipType`",
));
}
}
let all_unit = data
.variants
.iter()
.all(|variant| matches!(variant.fields, Fields::Unit));
let valid_union = data.variants.iter().all(|variant| {
matches!(variant.fields, Fields::Unit)
|| matches!(&variant.fields, Fields::Unnamed(fields) if fields.unnamed.len() == 1)
});
if all_unit {
derive_unit_enum(ident, type_name, type_documentation, data)
} else if valid_union {
derive_union_enum(ident, type_name, type_documentation, data)
} else {
let span = data
.variants
.iter()
.find(|variant| {
!matches!(variant.fields, Fields::Unit)
&& !matches!(&variant.fields, Fields::Unnamed(fields) if fields.unnamed.len() == 1)
})
.map_or_else(|| data.enum_token.span, Spanned::span);
Err(Error::new(
span,
"a `WipType` enum must contain unit variants and/or single-field tuple variants; named-field and multi-field variants are unsupported",
))
}
}
fn derive_unit_enum(
ident: &Ident,
type_name: &str,
type_documentation: Option<&Documentation>,
data: &DataEnum,
) -> syn::Result<TokenStream2> {
let declaration_doc = documentation_tokens(type_documentation);
let declarations = data.variants.iter().map(|variant| {
let name = source_name(&variant.ident);
let documentation = documentation_tokens(documentation(&variant.attrs).as_ref());
quote! {
::wip_interface::wip_protocol::EnumCase {
name: ::std::string::String::from(#name),
documentation: #documentation,
}
}
});
let encoders = data.variants.iter().map(|variant| {
let variant_ident = &variant.ident;
let name = source_name(variant_ident);
quote!(Self::#variant_ident => #name)
});
let decoders = data.variants.iter().map(|variant| {
let variant_ident = &variant.ident;
let name = source_name(variant_ident);
quote!(#name => ::std::result::Result::Ok(Self::#variant_ident))
});
Ok(quote! {
impl ::wip_interface::WipType for #ident {
fn type_expr() -> ::wip_interface::wip_protocol::TypeExpr {
::wip_interface::wip_protocol::TypeExpr::Named {
name: ::std::string::String::from(#type_name),
}
}
fn collect_declarations(
declarations: &mut ::std::vec::Vec<
::wip_interface::wip_protocol::TypeDeclaration,
>,
) {
let declaration = ::wip_interface::wip_protocol::TypeDeclaration {
name: ::std::string::String::from(#type_name),
documentation: #declaration_doc,
definition: ::wip_interface::wip_protocol::TypeExpr::Enum {
cases: ::std::vec![#(#declarations),*],
},
};
let _ = ::wip_interface::__private::register_declaration(
declarations,
declaration,
);
}
fn encode(
self,
) -> ::std::result::Result<
::wip_interface::wip_protocol::Value,
::wip_interface::CodecError,
> {
let variant = match self {
#(#encoders,)*
};
::std::result::Result::Ok(
::wip_interface::wip_protocol::Value::String(
::std::string::String::from(variant),
),
)
}
fn decode(
value: ::wip_interface::wip_protocol::Value,
) -> ::std::result::Result<Self, ::wip_interface::CodecError> {
match value {
::wip_interface::wip_protocol::Value::String(variant) => {
match variant.as_str() {
#(#decoders,)*
_ => ::std::result::Result::Err(
::wip_interface::__private::decode_unknown_variant(
#type_name,
variant,
),
),
}
}
actual => ::std::result::Result::Err(
::wip_interface::__private::decode_type_mismatch("enum", actual),
),
}
}
}
})
}
fn derive_union_enum(
ident: &Ident,
type_name: &str,
type_documentation: Option<&Documentation>,
data: &DataEnum,
) -> syn::Result<TokenStream2> {
struct Variant<'a> {
ident: &'a Ident,
name: String,
payload: Option<&'a Type>,
documentation: Option<Documentation>,
}
let mut variants = Vec::new();
for variant in &data.variants {
let payload = match &variant.fields {
Fields::Unit => None,
Fields::Unnamed(fields) if fields.unnamed.len() == 1 => {
let field = fields.unnamed.first().expect("length checked");
if direct_option(&field.ty)?.is_some() || contains_option(&field.ty) {
return Err(Error::new_spanned(
&field.ty,
"`Option<T>` is only supported directly on DTO fields and operation parameters",
));
}
validate_value_type(&field.ty, "union case payload")?;
Some(&field.ty)
}
_ => {
return Err(Error::new_spanned(
variant,
"a `WipType` union case must be unit-like or contain exactly one unnamed payload",
));
}
};
variants.push(Variant {
ident: &variant.ident,
name: source_name(&variant.ident),
payload,
documentation: documentation(&variant.attrs),
});
}
let declaration_doc = documentation_tokens(type_documentation);
let declarations = variants.iter().map(|variant| {
let name = &variant.name;
let documentation = documentation_tokens(variant.documentation.as_ref());
let payload = match variant.payload {
Some(ty) => quote! {
::std::option::Option::Some(
<#ty as ::wip_interface::WipType>::type_expr(),
)
},
None => quote!(::std::option::Option::None),
};
quote! {
::wip_interface::wip_protocol::UnionCase {
name: ::std::string::String::from(#name),
documentation: #documentation,
payload: #payload,
}
}
});
let dependency_collectors = variants.iter().filter_map(|variant| {
variant.payload.map(
|ty| quote!(<#ty as ::wip_interface::WipType>::collect_declarations(declarations);),
)
});
let encoders = variants.iter().map(|variant| {
let variant_ident = variant.ident;
let name = &variant.name;
match variant.payload {
Some(ty) => quote! {
Self::#variant_ident(value) => {
let mut fields = ::std::collections::BTreeMap::new();
fields.insert(
::std::string::String::from(
::wip_interface::wip_protocol::UNION_CASE_FIELD,
),
::wip_interface::wip_protocol::Value::String(
::std::string::String::from(#name),
),
);
fields.insert(
::std::string::String::from(
::wip_interface::wip_protocol::UNION_VALUE_FIELD,
),
<#ty as ::wip_interface::WipType>::encode(value)?,
);
::wip_interface::wip_protocol::Value::Record(fields)
}
},
None => quote! {
Self::#variant_ident => {
let mut fields = ::std::collections::BTreeMap::new();
fields.insert(
::std::string::String::from(
::wip_interface::wip_protocol::UNION_CASE_FIELD,
),
::wip_interface::wip_protocol::Value::String(
::std::string::String::from(#name),
),
);
::wip_interface::wip_protocol::Value::Record(fields)
}
},
}
});
let decoders = variants.iter().map(|variant| {
let variant_ident = variant.ident;
let name = &variant.name;
match variant.payload {
Some(ty) => quote! {
#name => {
let value = decoder.required::<#ty>(
::wip_interface::wip_protocol::UNION_VALUE_FIELD,
)?;
decoder.finish()?;
::std::result::Result::Ok(Self::#variant_ident(value))
}
},
None => quote! {
#name => {
decoder.finish()?;
::std::result::Result::Ok(Self::#variant_ident)
}
},
}
});
Ok(quote! {
impl ::wip_interface::WipType for #ident {
fn type_expr() -> ::wip_interface::wip_protocol::TypeExpr {
::wip_interface::wip_protocol::TypeExpr::Named {
name: ::std::string::String::from(#type_name),
}
}
fn collect_declarations(
declarations: &mut ::std::vec::Vec<
::wip_interface::wip_protocol::TypeDeclaration,
>,
) {
let declaration = ::wip_interface::wip_protocol::TypeDeclaration {
name: ::std::string::String::from(#type_name),
documentation: #declaration_doc,
definition: ::wip_interface::wip_protocol::TypeExpr::Union {
cases: ::std::vec![#(#declarations),*],
},
};
if ::wip_interface::__private::register_declaration(declarations, declaration) {
#(#dependency_collectors)*
}
}
fn encode(
self,
) -> ::std::result::Result<
::wip_interface::wip_protocol::Value,
::wip_interface::CodecError,
> {
let value = match self {
#(#encoders,)*
};
::std::result::Result::Ok(value)
}
fn decode(
value: ::wip_interface::wip_protocol::Value,
) -> ::std::result::Result<Self, ::wip_interface::CodecError> {
match value {
::wip_interface::wip_protocol::Value::Record(fields) => {
let mut decoder =
::wip_interface::__private::RecordDecoder::new(fields);
let case = decoder.required::<::std::string::String>(
::wip_interface::wip_protocol::UNION_CASE_FIELD,
)?;
match case.as_str() {
#(#decoders,)*
_ => ::std::result::Result::Err(
::wip_interface::__private::decode_unknown_variant(
#type_name,
case,
),
),
}
}
actual => ::std::result::Result::Err(
::wip_interface::__private::decode_type_mismatch("union record", actual),
),
}
}
}
})
}
fn optional_type(ty: &Type, context: &str) -> syn::Result<(bool, Type)> {
if let Some(inner) = direct_option(ty)? {
if contains_option(&inner) {
return Err(Error::new_spanned(
inner,
format!("nested `Option<T>` is not supported for {context}"),
));
}
Ok((true, inner))
} else {
if contains_option(ty) {
return Err(Error::new_spanned(
ty,
format!("`Option<T>` must appear directly as the {context} type"),
));
}
Ok((false, ty.clone()))
}
}
fn direct_option(ty: &Type) -> syn::Result<Option<Type>> {
let Type::Path(path) = ty else {
return Ok(None);
};
if path.qself.is_some() {
return Ok(None);
}
let Some(segment) = path.path.segments.last() else {
return Ok(None);
};
if segment.ident != "Option" {
return Ok(None);
}
let PathArguments::AngleBracketed(arguments) = &segment.arguments else {
return Err(Error::new_spanned(
segment,
"`Option` must have exactly one type argument",
));
};
if arguments.args.len() != 1 {
return Err(Error::new_spanned(
arguments,
"`Option` must have exactly one type argument",
));
}
match arguments.args.first() {
Some(GenericArgument::Type(inner)) => Ok(Some(inner.clone())),
_ => Err(Error::new_spanned(
arguments,
"`Option` must have exactly one type argument",
)),
}
}
fn validate_value_type(ty: &Type, context: &str) -> syn::Result<()> {
reject_custom_lifetimes(ty, context)?;
if let Some(span) = first_reference(ty) {
return Err(Error::new(
span,
format!("borrowed types are not supported as a WIP {context}"),
));
}
if let Some(span) = first_unit(ty) {
return Err(Error::new(
span,
format!(
"`()` is only supported as an operation application return, not as a {context}"
),
));
}
if contains_option(ty) {
return Err(Error::new_spanned(
ty,
"`Option<T>` is only supported directly on DTO fields and operation parameters",
));
}
Ok(())
}
fn reject_custom_lifetimes(ty: &Type, context: &str) -> syn::Result<()> {
struct Finder(Option<Span>);
impl<'ast> Visit<'ast> for Finder {
fn visit_lifetime(&mut self, lifetime: &'ast syn::Lifetime) {
if self.0.is_none() {
self.0 = Some(lifetime.span());
}
}
}
let mut finder = Finder(None);
finder.visit_type(ty);
if let Some(span) = finder.0 {
Err(Error::new(
span,
format!("custom lifetimes are not supported in a WIP {context}"),
))
} else {
Ok(())
}
}
fn contains_option(ty: &Type) -> bool {
struct Finder(bool);
impl<'ast> Visit<'ast> for Finder {
fn visit_type_path(&mut self, path: &'ast syn::TypePath) {
if path
.path
.segments
.last()
.is_some_and(|segment| segment.ident == "Option")
{
self.0 = true;
}
visit::visit_type_path(self, path);
}
}
let mut finder = Finder(false);
finder.visit_type(ty);
finder.0
}
fn first_reference(ty: &Type) -> Option<Span> {
struct Finder(Option<Span>);
impl<'ast> Visit<'ast> for Finder {
fn visit_type_reference(&mut self, reference: &'ast syn::TypeReference) {
if self.0.is_none() {
self.0 = Some(reference.span());
}
visit::visit_type_reference(self, reference);
}
}
let mut finder = Finder(None);
finder.visit_type(ty);
finder.0
}
fn first_unit(ty: &Type) -> Option<Span> {
struct Finder(Option<Span>);
impl<'ast> Visit<'ast> for Finder {
fn visit_type_tuple(&mut self, tuple: &'ast syn::TypeTuple) {
if tuple.elems.is_empty() && self.0.is_none() {
self.0 = Some(tuple.span());
}
visit::visit_type_tuple(self, tuple);
}
}
let mut finder = Finder(None);
finder.visit_type(ty);
finder.0
}
fn is_unit(ty: &Type) -> bool {
matches!(ty, Type::Tuple(tuple) if tuple.elems.is_empty())
}
fn source_name(ident: &Ident) -> String {
ident.unraw().to_string()
}
fn pascal_case(name: &str) -> String {
let mut result = String::new();
let mut uppercase = true;
for character in name.chars() {
if character == '_' {
uppercase = true;
} else if uppercase {
result.extend(character.to_uppercase());
uppercase = false;
} else {
result.push(character);
}
}
result
}
fn documentation(attributes: &[Attribute]) -> Option<Documentation> {
documentation_from_lines(doc_lines(attributes))
}
fn doc_lines(attributes: &[Attribute]) -> Vec<String> {
attributes
.iter()
.filter_map(|attribute| {
if !attribute.path().is_ident("doc") {
return None;
}
match &attribute.meta {
Meta::NameValue(name_value) => match &name_value.value {
syn::Expr::Lit(expression) => match &expression.lit {
Lit::Str(value) => {
let value = value.value();
Some(value.strip_prefix(' ').unwrap_or(&value).to_owned())
}
_ => None,
},
_ => None,
},
_ => None,
}
})
.collect()
}
fn split_returns_section(lines: Vec<String>) -> (Vec<String>, Vec<String>) {
let Some(start) = lines.iter().position(|line| line.trim() == "# Returns") else {
return (lines, Vec::new());
};
let end = lines[start + 1..]
.iter()
.position(|line| line.trim_start().starts_with("# "))
.map_or(lines.len(), |offset| start + 1 + offset);
let mut operation = lines[..start].to_vec();
operation.extend_from_slice(&lines[end..]);
(operation, lines[start + 1..end].to_vec())
}
fn documentation_from_lines(mut lines: Vec<String>) -> Option<Documentation> {
while lines.first().is_some_and(|line| line.trim().is_empty()) {
lines.remove(0);
}
while lines.last().is_some_and(|line| line.trim().is_empty()) {
lines.pop();
}
if lines.is_empty() {
return None;
}
let split = lines
.iter()
.position(|line| line.trim().is_empty())
.unwrap_or(lines.len());
let summary = lines[..split]
.iter()
.map(|line| line.trim())
.collect::<Vec<_>>()
.join(" ");
let mut details_lines = if split < lines.len() {
lines[split + 1..].to_vec()
} else {
Vec::new()
};
while details_lines
.first()
.is_some_and(|line| line.trim().is_empty())
{
details_lines.remove(0);
}
while details_lines
.last()
.is_some_and(|line| line.trim().is_empty())
{
details_lines.pop();
}
let details = (!details_lines.is_empty()).then(|| details_lines.join("\n"));
Some(Documentation { summary, details })
}
fn documentation_tokens(documentation: Option<&Documentation>) -> TokenStream2 {
match documentation {
Some(documentation) => {
let summary = &documentation.summary;
match &documentation.details {
Some(details) => quote! {
::wip_interface::__private::documentation(#summary, ::std::option::Option::Some(#details))
},
None => quote! {
::wip_interface::__private::documentation(#summary, ::std::option::Option::None)
},
}
}
None => quote! {
::wip_interface::__private::documentation("", ::std::option::Option::None)
},
}
}