extern crate proc_macro;
extern crate syn;
#[macro_use]
extern crate quote;
use concordium_contracts_common::*;
use proc_macro::TokenStream;
use proc_macro2::Span;
use quote::ToTokens;
#[cfg(feature = "build-schema")]
use std::collections::HashMap;
use std::{
collections::{BTreeMap, BTreeSet},
convert::TryFrom,
ops::Neg,
};
use syn::{
parse::Parser, parse_macro_input, punctuated::*, spanned::Spanned, DataEnum, Ident, Meta, Token,
};
fn unwrap_or_report(v: syn::Result<TokenStream>) -> TokenStream {
match v {
Ok(ts) => ts,
Err(e) => e.to_compile_error().into(),
}
}
fn attach_error<A>(mut v: syn::Result<A>, msg: &str) -> syn::Result<A> {
if let Err(e) = v.as_mut() {
let span = e.span();
e.combine(syn::Error::new(span, msg));
}
v
}
struct OptionalArguments {
pub(crate) payable: bool,
pub(crate) enable_logger: bool,
pub(crate) low_level: bool,
pub(crate) parameter: Option<syn::LitStr>,
}
struct InitAttributes {
pub(crate) contract: syn::LitStr,
pub(crate) optional: OptionalArguments,
}
struct ReceiveAttributes {
pub(crate) contract: syn::LitStr,
pub(crate) name: syn::LitStr,
pub(crate) optional: OptionalArguments,
}
#[derive(Default)]
struct ParsedAttributes {
pub(crate) flags: BTreeSet<syn::Ident>,
pub(crate) values: BTreeMap<syn::Ident, syn::LitStr>,
}
impl ParsedAttributes {
pub(crate) fn extract_value(&mut self, key: &str) -> Option<syn::LitStr> {
let key = syn::Ident::new(key, Span::call_site());
self.values.remove(&key)
}
pub(crate) fn extract_flag(&mut self, key: &str) -> bool {
let key = syn::Ident::new(key, Span::call_site());
self.flags.remove(&key)
}
pub(crate) fn report_all_attributes(self) -> syn::Result<()> {
let mut iter = self.flags.into_iter().chain(self.values.into_iter().map(|(k, _)| k));
if let Some(ident) = iter.next() {
let mut err =
syn::Error::new(ident.span(), format!("Unrecognized attribute {}.", ident));
for next_ident in iter {
err.combine(syn::Error::new(
ident.span(),
format!("Unrecognized attribute {}.", next_ident),
));
}
Err(err)
} else {
Ok(())
}
}
}
fn parse_attributes<'a>(iter: impl IntoIterator<Item = &'a Meta>) -> syn::Result<ParsedAttributes> {
let mut ret = ParsedAttributes::default();
let mut errors = Vec::new();
let mut duplicate_values = BTreeMap::new();
let mut duplicate_flags = BTreeMap::new();
for attr in iter.into_iter() {
match attr {
Meta::NameValue(mnv) => {
if let Some(ident) = mnv.path.get_ident() {
if let syn::Lit::Str(ls) = &mnv.lit {
if let Some((existing_ident, _)) = ret.values.get_key_value(ident) {
let v = duplicate_values.entry(ident).or_insert_with(|| {
syn::Error::new(
existing_ident.span(),
format!("Duplicate attribute '{}'.", existing_ident),
)
});
v.combine(syn::Error::new(
ident.span(),
format!("'{}' also appears here.", ident),
));
} else {
ret.values.insert(ident.clone(), ls.clone());
}
} else {
errors.push(syn::Error::new(
mnv.path.span(),
format!(
"Values of attribute must be string literals, e.g., '{} = \
\"value\"'",
ident
),
));
}
} else {
errors.push(syn::Error::new(
mnv.path.span(),
"Unrecognized attribute. Only attribute names consisting of a single \
identifier are recognized.",
))
}
}
Meta::Path(p) => {
if let Some(ident) = p.get_ident() {
if let Some(existing_ident) = ret.flags.get(ident) {
let v = duplicate_flags.entry(ident).or_insert_with(|| {
syn::Error::new(
existing_ident.span(),
format!("Duplicate attribute '{}'.", existing_ident),
)
});
v.combine(syn::Error::new(
ident.span(),
format!("'{}' also appears here.", ident),
));
} else {
ret.flags.insert(ident.clone());
}
} else {
errors.push(syn::Error::new(
p.span(),
"Unrecognized attribute. Only attribute names consisting of a single \
identifier are recognized.",
))
}
}
Meta::List(p) => {
errors.push(syn::Error::new(p.span(), "Unrecognized attribute."));
}
}
}
let mut iter = errors
.into_iter()
.chain(duplicate_values.into_iter().map(|(_, v)| v))
.chain(duplicate_flags.into_iter().map(|(_, v)| v));
if let Some(err) = iter.next() {
let mut err = err;
for next_err in iter {
err.combine(next_err);
}
Err(err)
} else {
Ok(ret)
}
}
#[cfg(feature = "build-schema")]
struct ContractStateAttributes {
pub(crate) contract: syn::LitStr,
}
#[cfg(feature = "build-schema")]
const CONTRACT_STATE_ATTRIBUTE_CONTRACT: &str = "contract";
#[cfg(feature = "build-schema")]
fn parse_contract_state_attributes<'a, I: IntoIterator<Item = &'a Meta>>(
attrs: I,
) -> syn::Result<ContractStateAttributes> {
let mut attributes = parse_attributes(attrs)?;
let contract =
attributes.extract_value(CONTRACT_STATE_ATTRIBUTE_CONTRACT).ok_or_else(|| {
syn::Error::new(
Span::call_site(),
"A name for the contract must be provided, using the 'contract' attribute.\n\nFor \
example, #[contract_state(contract = \"my-contract\")]",
)
})?;
attributes.report_all_attributes()?;
Ok(ContractStateAttributes {
contract,
})
}
const INIT_ATTRIBUTE_PARAMETER: &str = "parameter";
const INIT_ATTRIBUTE_CONTRACT: &str = "contract";
const INIT_ATTRIBUTE_PAYABLE: &str = "payable";
const INIT_ATTRIBUTE_ENABLE_LOGGER: &str = "enable_logger";
const INIT_ATTRIBUTE_LOW_LEVEL: &str = "low_level";
fn parse_init_attributes<'a, I: IntoIterator<Item = &'a Meta>>(
attrs: I,
) -> syn::Result<InitAttributes> {
let mut attributes = parse_attributes(attrs)?;
let contract: syn::LitStr =
attributes.extract_value(INIT_ATTRIBUTE_CONTRACT).ok_or_else(|| {
syn::Error::new(
Span::call_site(),
"A name for the contract must be provided, using the 'contract' attribute.\n\nFor \
example, #[init(contract = \"my-contract\")]",
)
})?;
let parameter: Option<syn::LitStr> = attributes.extract_value(INIT_ATTRIBUTE_PARAMETER);
let payable = attributes.extract_flag(INIT_ATTRIBUTE_PAYABLE);
let enable_logger = attributes.extract_flag(INIT_ATTRIBUTE_ENABLE_LOGGER);
let low_level = attributes.extract_flag(INIT_ATTRIBUTE_LOW_LEVEL);
attributes.report_all_attributes()?;
Ok(InitAttributes {
contract,
optional: OptionalArguments {
payable,
enable_logger,
low_level,
parameter,
},
})
}
const RECEIVE_ATTRIBUTE_PARAMETER: &str = "parameter";
const RECEIVE_ATTRIBUTE_CONTRACT: &str = "contract";
const RECEIVE_ATTRIBUTE_NAME: &str = "name";
const RECEIVE_ATTRIBUTE_PAYABLE: &str = "payable";
const RECEIVE_ATTRIBUTE_ENABLE_LOGGER: &str = "enable_logger";
const RECEIVE_ATTRIBUTE_LOW_LEVEL: &str = "low_level";
fn parse_receive_attributes<'a, I: IntoIterator<Item = &'a Meta>>(
attrs: I,
) -> syn::Result<ReceiveAttributes> {
let mut attributes = parse_attributes(attrs)?;
let contract = attributes.extract_value(RECEIVE_ATTRIBUTE_CONTRACT);
let name = attributes.extract_value(RECEIVE_ATTRIBUTE_NAME);
let parameter: Option<syn::LitStr> = attributes.extract_value(RECEIVE_ATTRIBUTE_PARAMETER);
let payable = attributes.extract_flag(RECEIVE_ATTRIBUTE_PAYABLE);
let enable_logger = attributes.extract_flag(RECEIVE_ATTRIBUTE_ENABLE_LOGGER);
let low_level = attributes.extract_flag(RECEIVE_ATTRIBUTE_LOW_LEVEL);
attributes.report_all_attributes()?;
match (contract, name) {
(Some(contract), Some(name)) => Ok(ReceiveAttributes {
contract,
name,
optional: OptionalArguments {
payable,
enable_logger,
low_level,
parameter,
},
}),
(Some(_), None) => Err(syn::Error::new(
Span::call_site(),
"A name for the method must be provided, using the 'name' attribute.\n\nFor example, \
#[receive(name = \"receive\")]",
)),
(None, Some(_)) => Err(syn::Error::new(
Span::call_site(),
"A name for the method must be provided, using the 'contract' attribute.\n\nFor \
example, #[receive(contract = \"my-contract\")]",
)),
(None, None) => Err(syn::Error::new(
Span::call_site(),
"A contract name and a name for the method must be provided, using the 'contract' and \
'name' attributes.\n\nFor example, #[receive(contract = \"my-contract\", name = \
\"receive\")]",
)),
}
}
fn contains_attribute<'a, I: IntoIterator<Item = &'a Meta>>(iter: I, name: &str) -> bool {
iter.into_iter().any(|attr| attr.path().is_ident(name))
}
#[proc_macro_attribute]
pub fn init(attr: TokenStream, item: TokenStream) -> TokenStream {
unwrap_or_report(init_worker(attr, item))
}
fn init_worker(attr: TokenStream, item: TokenStream) -> syn::Result<TokenStream> {
let ast: syn::ItemFn =
attach_error(syn::parse(item), "#[init] can only be applied to functions.")?;
let attrs = Punctuated::<Meta, Token![,]>::parse_terminated.parse(attr)?;
let init_attributes = parse_init_attributes(&attrs)?;
let contract_name = init_attributes.contract;
let fn_name = &ast.sig.ident;
let rust_export_fn_name = format_ident!("export_{}", fn_name);
let wasm_export_fn_name = format!("init_{}", contract_name.value());
if let Err(e) = ContractName::is_valid_contract_name(&wasm_export_fn_name) {
return Err(syn::Error::new(contract_name.span(), e));
}
let amount_ident = format_ident!("amount");
let mut required_args = vec!["ctx: &impl HasInitContext"];
let (setup_fn_optional_args, fn_optional_args) = contract_function_optional_args_tokens(
&init_attributes.optional,
&amount_ident,
&mut required_args,
);
let mut out = if init_attributes.optional.low_level {
required_args.push("state: &mut ContractState");
quote! {
#[export_name = #wasm_export_fn_name]
pub extern "C" fn #rust_export_fn_name(#amount_ident: concordium_std::Amount) -> i32 {
use concordium_std::{trap, ExternContext, InitContextExtern, ContractState};
#setup_fn_optional_args
let ctx = ExternContext::<InitContextExtern>::open(());
let mut state = ContractState::open(());
match #fn_name(&ctx, #(#fn_optional_args, )* &mut state) {
Ok(()) => 0,
Err(reject) => {
let code = Reject::from(reject).error_code.get();
if code < 0 {
code
} else {
trap() }
}
}
}
}
} else {
quote! {
#[export_name = #wasm_export_fn_name]
pub extern "C" fn #rust_export_fn_name(amount: concordium_std::Amount) -> i32 {
use concordium_std::{trap, ExternContext, InitContextExtern, ContractState};
#setup_fn_optional_args
let ctx = ExternContext::<InitContextExtern>::open(());
match #fn_name(&ctx, #(#fn_optional_args),*) {
Ok(state) => {
let mut state_bytes = ContractState::open(());
if state.serial(&mut state_bytes).is_err() {
trap() };
0
}
Err(reject) => {
let code = Reject::from(reject).error_code.get();
if code < 0 {
code
} else {
trap() }
}
}
}
}
};
let arg_count = ast.sig.inputs.len();
if arg_count != required_args.len() {
return Err(syn::Error::new(
ast.sig.inputs.span(),
format!(
"Incorrect number of function arguments, the expected arguments are ({}) ",
required_args.join(", ")
),
));
}
let parameter_option = init_attributes.optional.parameter.map(|a| a.value());
out.extend(contract_function_schema_tokens(
parameter_option,
rust_export_fn_name,
wasm_export_fn_name,
));
ast.to_tokens(&mut out);
Ok(out.into())
}
#[proc_macro_attribute]
pub fn receive(attr: TokenStream, item: TokenStream) -> TokenStream {
unwrap_or_report(receive_worker(attr, item))
}
fn receive_worker(attr: TokenStream, item: TokenStream) -> syn::Result<TokenStream> {
let ast: syn::ItemFn =
attach_error(syn::parse(item), "#[receive] can only be applied to functions.")?;
let attrs = Punctuated::<Meta, Token![,]>::parse_terminated.parse(attr)?;
let receive_attributes = parse_receive_attributes(&attrs)?;
let contract_name = receive_attributes.contract;
let method_name = receive_attributes.name;
let fn_name = &ast.sig.ident;
let rust_export_fn_name = format_ident!("export_{}", fn_name);
let wasm_export_fn_name = format!("{}.{}", contract_name.value(), method_name.value());
let contract_name_validation =
ContractName::is_valid_contract_name(&format!("init_{}", contract_name.value()))
.map_err(|e| syn::Error::new(contract_name.span(), e));
let receive_name_validation = ReceiveName::is_valid_receive_name(&wasm_export_fn_name)
.map_err(|e| syn::Error::new(method_name.span(), e));
match (contract_name_validation, receive_name_validation) {
(Err(mut e0), Err(e1)) => {
e0.combine(e1);
return Err(e0);
}
(Err(e), _) => return Err(e),
(_, Err(e)) => return Err(e),
_ => (),
};
let amount_ident = format_ident!("amount");
let mut required_args = vec!["ctx: &impl HasReceiveContext"];
let (setup_fn_optional_args, fn_optional_args) = contract_function_optional_args_tokens(
&receive_attributes.optional,
&amount_ident,
&mut required_args,
);
let mut out = if receive_attributes.optional.low_level {
required_args.push("state: &mut ContractState");
quote! {
#[export_name = #wasm_export_fn_name]
pub extern "C" fn #rust_export_fn_name(#amount_ident: concordium_std::Amount) -> i32 {
use concordium_std::{SeekFrom, ContractState, Logger, ReceiveContextExtern, ExternContext};
#setup_fn_optional_args
let ctx = ExternContext::<ReceiveContextExtern>::open(());
let mut state = ContractState::open(());
let res: Result<Action, _> = #fn_name(&ctx, #(#fn_optional_args, )* &mut state);
match res {
Ok(act) => {
act.tag() as i32
}
Err(reject) => {
let code = Reject::from(reject).error_code.get();
if code < 0 {
code
} else {
trap() }
}
}
}
}
} else {
required_args.push("state: &mut MyState");
quote! {
#[export_name = #wasm_export_fn_name]
pub extern "C" fn #rust_export_fn_name(#amount_ident: concordium_std::Amount) -> i32 {
use concordium_std::{SeekFrom, ContractState, Logger, trap};
#setup_fn_optional_args
let ctx = ExternContext::<ReceiveContextExtern>::open(());
let mut state_bytes = ContractState::open(());
if let Ok(mut state) = (&mut state_bytes).get() {
let res: Result<Action, _> = #fn_name(&ctx, #(#fn_optional_args, )* &mut state);
match res {
Ok(act) => {
let res = state_bytes
.seek(SeekFrom::Start(0))
.and_then(|_| state.serial(&mut state_bytes));
if res.is_err() {
trap() } else {
act.tag() as i32
}
}
Err(reject) => {
let code = Reject::from(reject).error_code.get();
if code < 0 {
code
} else {
trap() }
}
}
} else {
trap() }
}
}
};
let arg_count = ast.sig.inputs.len();
if arg_count != required_args.len() {
return Err(syn::Error::new(
ast.sig.inputs.span(),
format!(
"Incorrect number of function arguments, the expected arguments are ({}) ",
required_args.join(", ")
),
));
}
let parameter_option = receive_attributes.optional.parameter.map(|a| a.value());
out.extend(contract_function_schema_tokens(
parameter_option,
rust_export_fn_name,
wasm_export_fn_name,
));
ast.to_tokens(&mut out);
Ok(out.into())
}
fn contract_function_optional_args_tokens(
optional: &OptionalArguments,
amount_ident: &syn::Ident,
required_args: &mut Vec<&str>,
) -> (proc_macro2::TokenStream, Vec<proc_macro2::TokenStream>) {
let mut setup_fn_args = proc_macro2::TokenStream::new();
let mut fn_args = vec![];
if optional.payable {
required_args.push("amount: Amount");
fn_args.push(quote!(#amount_ident));
} else {
setup_fn_args.extend(quote! {
if #amount_ident.micro_ccd != 0 {
return concordium_std::Reject::from(concordium_std::NotPayableError).error_code.get();
}
});
};
if optional.enable_logger {
required_args.push("logger: &mut impl HasLogger");
let logger_ident = format_ident!("logger");
setup_fn_args.extend(quote!(let mut #logger_ident = concordium_std::Logger::init();));
fn_args.push(quote!(&mut #logger_ident));
}
(setup_fn_args, fn_args)
}
#[cfg(feature = "build-schema")]
fn contract_function_schema_tokens(
parameter_option: Option<String>,
rust_name: syn::Ident,
wasm_name: String,
) -> proc_macro2::TokenStream {
match parameter_option {
Some(parameter_ty) => {
let parameter_ident = syn::Ident::new(¶meter_ty, Span::call_site());
let schema_name = format!("concordium_schema_function_{}", wasm_name);
let schema_ident = format_ident!("concordium_schema_function_{}", rust_name);
quote! {
#[export_name = #schema_name]
pub extern "C" fn #schema_ident() -> *mut u8 {
let schema = <#parameter_ident as schema::SchemaType>::get_type();
let schema_bytes = concordium_std::to_bytes(&schema);
concordium_std::put_in_memory(&schema_bytes)
}
}
}
None => proc_macro2::TokenStream::new(),
}
}
#[cfg(not(feature = "build-schema"))]
fn contract_function_schema_tokens(
_parameter_option: Option<String>,
_rust_name: syn::Ident,
_wasm_name: String,
) -> proc_macro2::TokenStream {
proc_macro2::TokenStream::new()
}
#[proc_macro_derive(Deserial, attributes(concordium))]
pub fn deserial_derive(input: TokenStream) -> TokenStream {
let ast = parse_macro_input!(input);
unwrap_or_report(impl_deserial(&ast))
}
const CONCORDIUM_FIELD_ATTRIBUTE: &str = "concordium";
const VALID_CONCORDIUM_FIELD_ATTRIBUTES: [&str; 3] = ["size_length", "ensure_ordered", "rename"];
fn get_concordium_field_attributes(attributes: &[syn::Attribute]) -> syn::Result<Vec<syn::Meta>> {
attributes
.iter()
.flat_map(|attr| match attr.parse_meta() {
Ok(syn::Meta::List(list)) if list.path.is_ident(CONCORDIUM_FIELD_ATTRIBUTE) => {
list.nested
}
_ => syn::punctuated::Punctuated::new(),
})
.map(|nested| match nested {
syn::NestedMeta::Meta(meta) => {
let path = meta.path();
if VALID_CONCORDIUM_FIELD_ATTRIBUTES.iter().any(|&attr| path.is_ident(attr)) {
Ok(meta)
} else {
Err(syn::Error::new(meta.span(),
format!("The attribute '{}' is not supported as a concordium field attribute.",
path.to_token_stream())
))
}
}
lit => Err(syn::Error::new(lit.span(), "Literals are not supported in a concordium field attribute.")),
})
.collect()
}
fn find_field_attribute_value(
attributes: &[syn::Attribute],
target_attr: &str,
) -> syn::Result<Option<syn::Lit>> {
let target_attr = format_ident!("{}", target_attr);
let attr_values: Vec<_> = get_concordium_field_attributes(attributes)?
.into_iter()
.filter_map(|nested_meta| match nested_meta {
syn::Meta::NameValue(value) if value.path.is_ident(&target_attr) => Some(value.lit),
_ => None,
})
.collect();
if attr_values.is_empty() {
return Ok(None);
}
if attr_values.len() > 1 {
let mut init_error = syn::Error::new(
attr_values[1].span(),
format!("Attribute '{}' should only be specified once.", target_attr),
);
for other in attr_values.iter().skip(2) {
init_error.combine(syn::Error::new(
other.span(),
format!("Attribute '{}' should only be specified once.", target_attr),
))
}
Err(init_error)
} else {
Ok(Some(attr_values[0].clone()))
}
}
fn find_length_attribute(attributes: &[syn::Attribute]) -> syn::Result<Option<u32>> {
let value = match find_field_attribute_value(attributes, "size_length")? {
Some(v) => v,
None => return Ok(None),
};
let value_span = value.span();
let value = match value {
syn::Lit::Int(int) => int,
_ => return Err(syn::Error::new(value_span, "Length attribute value must be an integer.")),
};
let value = match value.base10_parse() {
Ok(v) => v,
_ => {
return Err(syn::Error::new(
value_span,
"Length attribute value must be a base 10 integer.",
))
}
};
match value {
1 | 2 | 4 | 8 => Ok(Some(value)),
_ => Err(syn::Error::new(value_span, "Length info must be either 1, 2, 4, or 8.")),
}
}
#[cfg(feature = "build-schema")]
fn find_rename_attribute(attributes: &[syn::Attribute]) -> syn::Result<Option<(String, Span)>> {
let value = match find_field_attribute_value(attributes, "rename")? {
Some(v) => v,
None => return Ok(None),
};
match value {
syn::Lit::Str(value) => Ok(Some((value.value(), value.span()))),
_ => Err(syn::Error::new(value.span(), "Rename attribute value must be a string.")),
}
}
#[cfg(feature = "build-schema")]
fn check_for_name_collisions(
used_names: &mut HashMap<String, Span>,
new_name: &str,
new_span: Span,
) -> syn::Result<()> {
if let Some(used_span) = used_names.insert(String::from(new_name), new_span) {
let error_msg = format!("the name `{}` is defined multiple times", new_name);
let mut error_at_first_def = syn::Error::new(used_span, &error_msg);
let error_at_second_def = syn::Error::new(new_span, &error_msg);
error_at_first_def.combine(error_at_second_def);
return Err(error_at_first_def);
}
Ok(())
}
fn impl_deserial_field(
f: &syn::Field,
ident: &syn::Ident,
source: &syn::Ident,
) -> syn::Result<proc_macro2::TokenStream> {
let concordium_attributes = get_concordium_field_attributes(&f.attrs)?;
let ensure_ordered = contains_attribute(&concordium_attributes, "ensure_ordered");
let size_length = find_length_attribute(&f.attrs)?;
let has_ctx = ensure_ordered || size_length.is_some();
let ty = &f.ty;
if has_ctx {
let l = format_ident!("U{}", 8 * size_length.unwrap_or(4));
Ok(quote! {
let #ident = <#ty as concordium_std::DeserialCtx>::deserial_ctx(concordium_std::schema::SizeLength::#l, #ensure_ordered, #source)?;
})
} else {
Ok(quote! {
let #ident = <#ty as Deserial>::deserial(#source)?;
})
}
}
fn impl_deserial(ast: &syn::DeriveInput) -> syn::Result<TokenStream> {
let data_name = &ast.ident;
let span = ast.span();
let read_ident = format_ident!("__R", span = span);
let (impl_generics, ty_generics, where_clauses) = ast.generics.split_for_impl();
let source_ident = Ident::new("source", Span::call_site());
let body_tokens = match ast.data {
syn::Data::Struct(ref data) => {
let mut names = proc_macro2::TokenStream::new();
let mut field_tokens = proc_macro2::TokenStream::new();
let return_tokens = match data.fields {
syn::Fields::Named(_) => {
for field in data.fields.iter() {
let field_ident = field.ident.clone().unwrap(); field_tokens.extend(impl_deserial_field(
field,
&field_ident,
&source_ident,
));
names.extend(quote!(#field_ident,))
}
quote!(Ok(#data_name{#names}))
}
syn::Fields::Unnamed(_) => {
for (i, f) in data.fields.iter().enumerate() {
let field_ident = format_ident!("x_{}", i);
field_tokens.extend(impl_deserial_field(f, &field_ident, &source_ident));
names.extend(quote!(#field_ident,))
}
quote!(Ok(#data_name(#names)))
}
_ => quote!(Ok(#data_name{})),
};
quote! {
#field_tokens
#return_tokens
}
}
syn::Data::Enum(ref data) => {
let mut matches_tokens = proc_macro2::TokenStream::new();
let source = Ident::new("source", Span::call_site());
let size = if data.variants.len() <= 256 {
format_ident!("u8")
} else if data.variants.len() <= 256 * 256 {
format_ident!("u16")
} else {
return Err(syn::Error::new(
ast.span(),
"[derive(Deserial)]: Too many variants. Maximum 65536 are supported.",
));
};
for (i, variant) in data.variants.iter().enumerate() {
let (field_names, pattern) = match variant.fields {
syn::Fields::Named(_) => {
let field_names: Vec<_> = variant
.fields
.iter()
.map(|field| field.ident.clone().unwrap())
.collect();
(field_names.clone(), quote! { {#(#field_names),*} })
}
syn::Fields::Unnamed(_) => {
let field_names: Vec<_> = variant
.fields
.iter()
.enumerate()
.map(|(i, _)| format_ident!("x_{}", i))
.collect();
(field_names.clone(), quote! { ( #(#field_names),* ) })
}
syn::Fields::Unit => (Vec::new(), proc_macro2::TokenStream::new()),
};
let field_tokens: proc_macro2::TokenStream = field_names
.iter()
.zip(variant.fields.iter())
.map(|(name, field)| impl_deserial_field(field, name, &source))
.collect::<syn::Result<proc_macro2::TokenStream>>()?;
let idx_lit = syn::LitInt::new(i.to_string().as_str(), Span::call_site());
let variant_ident = &variant.ident;
matches_tokens.extend(quote! {
#idx_lit => {
#field_tokens
Ok(#data_name::#variant_ident#pattern)
},
})
}
quote! {
let idx = #size::deserial(#source)?;
match idx {
#matches_tokens
_ => Err(Default::default())
}
}
}
_ => unimplemented!("#[derive(Deserial)] is not implemented for union."),
};
let gen = quote! {
#[automatically_derived]
impl #impl_generics Deserial for #data_name #ty_generics #where_clauses {
fn deserial<#read_ident: Read>(#source_ident: &mut #read_ident) -> ParseResult<Self> {
#body_tokens
}
}
};
Ok(gen.into())
}
#[proc_macro_derive(Serial, attributes(concordium))]
pub fn serial_derive(input: TokenStream) -> TokenStream {
let ast = parse_macro_input!(input);
unwrap_or_report(impl_serial(&ast))
}
fn impl_serial_field(
field: &syn::Field,
ident: &proc_macro2::TokenStream,
out: &syn::Ident,
) -> syn::Result<proc_macro2::TokenStream> {
if let Some(size_length) = find_length_attribute(&field.attrs)? {
let l = format_ident!("U{}", 8 * size_length);
Ok(quote!({
use concordium_std::SerialCtx;
#ident.serial_ctx(concordium_std::schema::SizeLength::#l, #out)?;
}))
} else {
Ok(quote! {
#ident.serial(#out)?;
})
}
}
fn impl_serial(ast: &syn::DeriveInput) -> syn::Result<TokenStream> {
let data_name = &ast.ident;
let span = ast.span();
let write_ident = format_ident!("W", span = span);
let (impl_generics, ty_generics, where_clauses) = ast.generics.split_for_impl();
let out_ident = format_ident!("out");
let body = match ast.data {
syn::Data::Struct(ref data) => {
let fields_tokens = match data.fields {
syn::Fields::Named(_) => {
data.fields
.iter()
.map(|field| {
let field_ident = field.ident.clone().unwrap(); let field_ident = quote!(self.#field_ident);
impl_serial_field(field, &field_ident, &out_ident)
})
.collect::<syn::Result<_>>()?
}
syn::Fields::Unnamed(_) => data
.fields
.iter()
.enumerate()
.map(|(i, field)| {
let i = syn::LitInt::new(i.to_string().as_str(), Span::call_site());
let field_ident = quote!(self.#i);
impl_serial_field(field, &field_ident, &out_ident)
})
.collect::<syn::Result<_>>()?,
syn::Fields::Unit => proc_macro2::TokenStream::new(),
};
quote! {
#fields_tokens
Ok(())
}
}
syn::Data::Enum(ref data) => {
let mut matches_tokens = proc_macro2::TokenStream::new();
let size = if data.variants.len() <= 256 {
format_ident!("u8")
} else if data.variants.len() <= 256 * 256 {
format_ident!("u16")
} else {
unimplemented!(
"[derive(Serial)]: Enums with more than 65536 variants are not supported."
);
};
for (i, variant) in data.variants.iter().enumerate() {
let (field_names, pattern) = match variant.fields {
syn::Fields::Named(_) => {
let field_names: Vec<_> = variant
.fields
.iter()
.map(|field| field.ident.clone().unwrap())
.collect();
(field_names.clone(), quote! { {#(#field_names),*} })
}
syn::Fields::Unnamed(_) => {
let field_names: Vec<_> = variant
.fields
.iter()
.enumerate()
.map(|(i, _)| format_ident!("x_{}", i))
.collect();
(field_names.clone(), quote! { (#(#field_names),*) })
}
syn::Fields::Unit => (Vec::new(), proc_macro2::TokenStream::new()),
};
let field_tokens: proc_macro2::TokenStream = field_names
.iter()
.zip(variant.fields.iter())
.map(|(name, field)| impl_serial_field(field, "e!(#name), &out_ident))
.collect::<syn::Result<_>>()?;
let idx_lit =
syn::LitInt::new(format!("{}{}", i, size).as_str(), Span::call_site());
let variant_ident = &variant.ident;
matches_tokens.extend(quote! {
#data_name::#variant_ident#pattern => {
#idx_lit.serial(#out_ident)?;
#field_tokens
},
})
}
quote! {
match self {
#matches_tokens
}
Ok(())
}
}
_ => unimplemented!("#[derive(Serial)] is not implemented for union."),
};
let gen = quote! {
#[automatically_derived]
impl #impl_generics Serial for #data_name #ty_generics #where_clauses {
fn serial<#write_ident: Write>(&self, #out_ident: &mut #write_ident) -> Result<(), #write_ident::Err> {
#body
}
}
};
Ok(gen.into())
}
#[proc_macro_derive(Serialize, attributes(concordium))]
pub fn serialize_derive(input: TokenStream) -> TokenStream {
unwrap_or_report(serialize_derive_worker(input))
}
fn serialize_derive_worker(input: TokenStream) -> syn::Result<TokenStream> {
let ast = syn::parse(input)?;
let mut tokens = impl_deserial(&ast)?;
tokens.extend(impl_serial(&ast)?);
Ok(tokens)
}
#[proc_macro_attribute]
pub fn contract_state(attr: TokenStream, item: TokenStream) -> TokenStream {
unwrap_or_report(contract_state_worker(attr, item))
}
#[cfg(feature = "build-schema")]
fn contract_state_worker(attr: TokenStream, item: TokenStream) -> syn::Result<TokenStream> {
let mut out = proc_macro2::TokenStream::new();
let data_ident = if let Ok(ast) = syn::parse::<syn::ItemStruct>(item.clone()) {
ast.to_tokens(&mut out);
ast.ident
} else if let Ok(ast) = syn::parse::<syn::ItemEnum>(item.clone()) {
ast.to_tokens(&mut out);
ast.ident
} else if let Ok(ast) = syn::parse::<syn::ItemType>(item.clone()) {
ast.to_tokens(&mut out);
ast.ident
} else {
return Err(syn::Error::new_spanned(
proc_macro2::TokenStream::from(item),
"#[contract_state] only supports structs, enums and type aliases.",
));
};
let attrs = Punctuated::<Meta, Token![,]>::parse_terminated.parse(attr)?;
let contract_state_attributes = parse_contract_state_attributes(&attrs)?;
let contract_name = contract_state_attributes.contract;
let wasm_schema_name = format!("concordium_schema_state_{}", contract_name.value());
let rust_schema_name = format_ident!("concordium_schema_state_{}", data_ident);
let generate_schema_tokens = quote! {
#[allow(non_snake_case)]
#[export_name = #wasm_schema_name]
pub extern "C" fn #rust_schema_name() -> *mut u8 {
let schema = <#data_ident as concordium_std::schema::SchemaType>::get_type();
let schema_bytes = concordium_std::to_bytes(&schema);
concordium_std::put_in_memory(&schema_bytes)
}
};
generate_schema_tokens.to_tokens(&mut out);
Ok(out.into())
}
#[cfg(not(feature = "build-schema"))]
fn contract_state_worker(_attr: TokenStream, item: TokenStream) -> syn::Result<TokenStream> {
Ok(item)
}
#[proc_macro_derive(SchemaType, attributes(size_length))]
pub fn schema_type_derive(input: TokenStream) -> TokenStream {
unwrap_or_report(schema_type_derive_worker(input))
}
#[cfg(feature = "build-schema")]
fn schema_type_derive_worker(input: TokenStream) -> syn::Result<TokenStream> {
let ast: syn::DeriveInput = syn::parse(input)?;
let data_name = &ast.ident;
let (impl_generics, ty_generics, where_clauses) = ast.generics.split_for_impl();
let body = match ast.data {
syn::Data::Struct(ref data) => {
let fields_tokens = schema_type_fields(&data.fields)?;
quote! {
concordium_std::schema::Type::Struct(#fields_tokens)
}
}
syn::Data::Enum(ref data) => {
let mut used_variant_names = HashMap::new();
let variant_tokens: Vec<_> = data
.variants
.iter()
.map(|variant| {
let (variant_name, variant_span) = match find_rename_attribute(&variant.attrs)?
{
Some(name_and_span) => name_and_span,
None => (variant.ident.to_string(), variant.ident.span()),
};
check_for_name_collisions(
&mut used_variant_names,
&variant_name,
variant_span,
)?;
let fields_tokens = schema_type_fields(&variant.fields)?;
Ok(quote! {
(concordium_std::String::from(#variant_name), #fields_tokens)
})
})
.collect::<syn::Result<_>>()?;
quote! {
concordium_std::schema::Type::Enum(concordium_std::Vec::from([ #(#variant_tokens),* ]))
}
}
_ => syn::Error::new(ast.span(), "Union is not supported").to_compile_error(),
};
let out = quote! {
#[automatically_derived]
impl #impl_generics concordium_std::schema::SchemaType for #data_name #ty_generics #where_clauses {
fn get_type() -> concordium_std::schema::Type {
#body
}
}
};
Ok(out.into())
}
#[cfg(not(feature = "build-schema"))]
fn schema_type_derive_worker(_input: TokenStream) -> syn::Result<TokenStream> {
Ok(TokenStream::new())
}
#[cfg(feature = "build-schema")]
fn schema_type_field_type(field: &syn::Field) -> syn::Result<proc_macro2::TokenStream> {
let field_type = &field.ty;
if let Some(l) = find_length_attribute(&field.attrs)? {
let size = format_ident!("U{}", 8 * l);
Ok(quote! {
<#field_type as concordium_std::schema::SchemaType>::get_type().set_size_length(concordium_std::schema::SizeLength::#size)
})
} else {
Ok(quote! {
<#field_type as concordium_std::schema::SchemaType>::get_type()
})
}
}
#[cfg(feature = "build-schema")]
fn schema_type_fields(fields: &syn::Fields) -> syn::Result<proc_macro2::TokenStream> {
match fields {
syn::Fields::Named(_) => {
let mut used_field_names = HashMap::new();
let fields_tokens: Vec<_> = fields
.iter()
.map(|field| {
let (field_name, field_span) = match find_rename_attribute(&field.attrs)? {
Some(name_and_span) => name_and_span,
None => (field.ident.clone().unwrap().to_string(), field.ident.span()), };
check_for_name_collisions(&mut used_field_names, &field_name, field_span)?;
let field_schema_type = schema_type_field_type(&field)?;
Ok(quote! {
(concordium_std::String::from(#field_name), #field_schema_type)
})
})
.collect::<syn::Result<_>>()?;
Ok(
quote! { concordium_std::schema::Fields::Named(concordium_std::Vec::from([ #(#fields_tokens),* ])) },
)
}
syn::Fields::Unnamed(_) => {
let fields_tokens: Vec<_> =
fields.iter().map(schema_type_field_type).collect::<syn::Result<_>>()?;
Ok(quote! { concordium_std::schema::Fields::Unnamed([ #(#fields_tokens),* ].to_vec()) })
}
syn::Fields::Unit => Ok(quote! { concordium_std::schema::Fields::None }),
}
}
const RESERVED_ERROR_CODES: i32 = i32::MIN + 100;
#[proc_macro_derive(Reject, attributes(from))]
pub fn reject_derive(input: TokenStream) -> TokenStream {
unwrap_or_report(reject_derive_worker(input))
}
fn reject_derive_worker(input: TokenStream) -> syn::Result<TokenStream> {
let ast: syn::DeriveInput = syn::parse(input)?;
let enum_data = match &ast.data {
syn::Data::Enum(data) => Ok(data),
_ => Err(syn::Error::new(ast.span(), "Reject can only be derived for enums.")),
}?;
let enum_ident = &ast.ident;
let too_many_variants = format!(
"Error enum {} cannot have more than {} variants.",
enum_ident,
RESERVED_ERROR_CODES.neg()
);
match i32::try_from(enum_data.variants.len()) {
Ok(n) if n <= RESERVED_ERROR_CODES.neg() => (),
_ => {
return Err(syn::Error::new(ast.span(), &too_many_variants));
}
};
let variant_error_conversions = generate_variant_error_conversions(&enum_data, &enum_ident)?;
let gen = quote! {
#[automatically_derived]
impl From<#enum_ident> for Reject {
#[inline(always)]
fn from(e: #enum_ident) -> Self {
Reject { error_code: unsafe { concordium_std::num::NonZeroI32::new_unchecked(-(e as i32) - 1) } }
}
}
#(#variant_error_conversions)*
};
Ok(gen.into())
}
fn generate_variant_error_conversions(
enum_data: &DataEnum,
enum_name: &syn::Ident,
) -> syn::Result<Vec<proc_macro2::TokenStream>> {
Ok(enum_data
.variants
.iter()
.map(|variant| {
if let Some((_, discriminant)) = variant.discriminant.as_ref() {
return Err(syn::Error::new(
discriminant.span(),
"Explicit discriminants are not yet supported.",
));
}
let variant_attributes = variant.attrs.iter();
variant_attributes
.map(move |attr| {
parse_attr_and_gen_error_conversions(attr, enum_name, &variant.ident)
})
.collect::<syn::Result<Vec<_>>>()
})
.collect::<syn::Result<Vec<_>>>()?
.into_iter()
.flatten()
.flatten()
.collect())
}
fn parse_attr_and_gen_error_conversions(
attr: &syn::Attribute,
enum_name: &syn::Ident,
variant_name: &syn::Ident,
) -> syn::Result<Vec<proc_macro2::TokenStream>> {
let wrong_from_usage = |x: &dyn Spanned| {
syn::Error::new(
x.span(),
"The `from` attribute expects a list of error types, e.g.: #[from(ParseError)].",
)
};
match attr.parse_meta() {
Ok(syn::Meta::List(list)) if list.path.is_ident("from") => {
let mut from_error_names = vec![];
for nested in list.nested.iter() {
match nested {
syn::NestedMeta::Meta(meta) => match meta {
Meta::Path(from_error) => {
let ident = from_error
.get_ident()
.ok_or_else(|| wrong_from_usage(from_error))?;
from_error_names.push(ident);
}
other => return Err(wrong_from_usage(&other)),
},
syn::NestedMeta::Lit(l) => return Err(wrong_from_usage(&l)),
}
}
Ok(from_error_token_stream(&from_error_names, &enum_name, variant_name).collect())
}
Ok(syn::Meta::NameValue(mnv)) if mnv.path.is_ident("from") => Err(wrong_from_usage(&mnv)),
_ => Ok(vec![]),
}
}
fn from_error_token_stream<'a>(
paths: &'a [&'a syn::Ident],
enum_name: &'a syn::Ident,
variant_name: &'a syn::Ident,
) -> impl Iterator<Item = proc_macro2::TokenStream> + 'a {
paths.iter().map(move |from_error| {
quote! {
impl From<#from_error> for #enum_name {
#[inline]
fn from(fe: #from_error) -> Self {
#enum_name::#variant_name
}
}}
})
}
#[proc_macro_attribute]
pub fn concordium_test(attr: TokenStream, item: TokenStream) -> TokenStream {
unwrap_or_report(concordium_test_worker(attr, item))
}
#[cfg(feature = "wasm-test")]
fn concordium_test_worker(_attr: TokenStream, item: TokenStream) -> syn::Result<TokenStream> {
let test_fn_ast: syn::ItemFn =
attach_error(syn::parse(item), "#[concordium_test] can only be applied to functions.")?;
let test_fn_name = &test_fn_ast.sig.ident;
let rust_export_fn_name = format_ident!("concordium_test_{}", test_fn_name);
let wasm_export_fn_name = format!("concordium_test {}", test_fn_name);
let test_fn = quote! {
#test_fn_ast
#[export_name = #wasm_export_fn_name]
pub extern "C" fn #rust_export_fn_name() {
#test_fn_name()
}
};
Ok(test_fn.into())
}
#[cfg(not(feature = "wasm-test"))]
fn concordium_test_worker(_attr: TokenStream, item: TokenStream) -> syn::Result<TokenStream> {
let test_fn_ast: syn::ItemFn =
attach_error(syn::parse(item), "#[concordium_test] can only be applied to functions.")?;
let test_fn = quote! {
#[test]
#test_fn_ast
};
Ok(test_fn.into())
}
#[cfg(feature = "wasm-test")]
#[proc_macro_attribute]
pub fn concordium_cfg_test(_attr: TokenStream, item: TokenStream) -> TokenStream { item }
#[cfg(not(feature = "wasm-test"))]
#[proc_macro_attribute]
pub fn concordium_cfg_test(_attr: TokenStream, item: TokenStream) -> TokenStream {
let item = proc_macro2::TokenStream::from(item);
let out = quote! {
#[cfg(test)]
#item
};
out.into()
}