use cfg_if::cfg_if;
use convert_case::{Case, Casing};
use proc_macro::TokenStream;
use proc_macro_error::emit_error;
use quote::{quote, ToTokens};
use syn::{parse_macro_input, parse_quote, spanned::Spanned};
use types::{
FnArgExtension, FnExtension, FnKind, InterfaceExtension, PublicFn, PublicFnArg, PublicImpl,
PublicTrait,
};
use crate::{
types::Purity,
utils::{
attrs::{check_attr_is_empty, consume_attr, consume_flag},
get_generics,
},
};
mod attrs;
mod types;
cfg_if! {
if #[cfg(feature = "export-abi")] {
mod export_abi;
type Extension = export_abi::InterfaceAbi;
} else {
type Extension = ();
}
}
const STYLUS_PUBLIC_TAG_CHECK_FN_NAME: &str =
"__stylus_trait_and_impl_must_be_tagged_with_public_macro";
pub fn public(attr: TokenStream, input: TokenStream) -> TokenStream {
check_attr_is_empty(attr);
let mut output = proc_macro2::TokenStream::new();
let item = parse_macro_input!(input as syn::Item);
match item {
syn::Item::Impl(mut item_impl) => {
let public_impl = PublicImpl::<Extension>::from(&mut item_impl);
add_stylus_public_tag_check_fn_definition(&mut item_impl);
output.extend(quote! {
#[cfg(not(feature = "contract-client-gen"))]
#[allow(dead_code)]
});
output.extend(item_impl.into_token_stream());
public_impl.to_tokens(&mut output);
}
syn::Item::Trait(mut item_trait) => {
let public_trait = PublicTrait::from(&mut item_trait);
add_stylus_public_tag_check_fn_declaration(&mut item_trait);
output.extend(quote! {
#[cfg(not(feature = "contract-client-gen"))]
#[allow(dead_code)]
});
output.extend(item_trait.into_token_stream());
public_trait.to_tokens(&mut output);
}
_ => {
emit_error!(item.span(), "expected impl or trait");
}
}
output.into()
}
fn add_stylus_public_tag_check_fn_declaration(item_trait: &mut syn::ItemTrait) {
let fn_name = syn::Ident::new(STYLUS_PUBLIC_TAG_CHECK_FN_NAME, item_trait.span());
let item: syn::TraitItem = parse_quote! {
fn #fn_name(&self);
};
item_trait.items.push(item);
}
fn add_stylus_public_tag_check_fn_definition(item_impl: &mut syn::ItemImpl) {
let fn_name = syn::Ident::new(STYLUS_PUBLIC_TAG_CHECK_FN_NAME, item_impl.span());
let item: syn::ImplItem = parse_quote! {
fn #fn_name(&self) {
}
};
item_impl.items.push(item);
}
impl From<&mut syn::ItemTrait> for PublicTrait {
fn from(node: &mut syn::ItemTrait) -> Self {
let ident = node.ident.clone();
let funcs = node
.items
.iter_mut()
.filter_map(|item| match item {
syn::TraitItem::Fn(func) => Some(PublicFn::from(func)),
syn::TraitItem::Const(_) => {
emit_error!(item, "unsupported trait item");
None
}
_ => {
None
}
})
.collect();
let (generic_params, where_clause) = get_generics(&node.generics);
let mut associated_types = Vec::new();
for item in &node.items {
if let syn::TraitItem::Type(type_item) = item {
associated_types.push((type_item.ident.clone(), type_item.bounds.clone()));
}
}
Self {
ident,
generic_params,
where_clause,
funcs,
associated_types,
}
}
}
impl ToTokens for PublicTrait {
fn to_tokens(&self, tokens: &mut proc_macro2::TokenStream) {
tokens.extend(self.contract_client_gen());
if self.generic_params.is_empty() {
for check in types::selector_collision_checks(&self.funcs) {
check.to_tokens(tokens);
}
}
}
}
impl ToTokens for PublicImpl {
fn to_tokens(&self, tokens: &mut proc_macro2::TokenStream) {
tokens.extend(self.struct_for_export_abi());
tokens.extend(self.contract_client_gen());
tokens.extend(self.print_from_args_fn());
if self.generic_params.is_empty() {
for check in types::selector_collision_checks(&self.funcs) {
check.to_tokens(tokens);
}
}
self.impl_router().to_tokens(tokens);
Extension::codegen(self).to_tokens(tokens);
}
}
impl From<&mut syn::ItemImpl> for PublicImpl {
fn from(node: &mut syn::ItemImpl) -> Self {
let mut implements = Vec::new();
if let Some(attr) = consume_attr::<attrs::Implements>(&mut node.attrs, "implements") {
implements.extend(attr.types);
}
let funcs = node
.items
.iter_mut()
.filter_map(|item| match item {
syn::ImplItem::Fn(func) => Some(PublicFn::from(func)),
syn::ImplItem::Const(_) => {
emit_error!(item, "unsupported impl item");
None
}
_ => {
None
}
})
.collect();
let self_ty = (*node.self_ty).clone();
let (generic_params, where_clause) = get_generics(&node.generics);
let trait_ = match &node.trait_ {
Some((_, trait_, _)) => Some(trait_.clone()),
_ => None,
};
let mut associated_types = Vec::new();
for item in &node.items {
if let syn::ImplItem::Type(type_item) = item {
associated_types.push((type_item.ident.clone(), type_item.ty.clone()));
}
}
#[allow(clippy::let_unit_value)]
let extension = <Extension as InterfaceExtension>::build(node);
Self {
self_ty,
generic_params,
where_clause,
trait_,
implements,
funcs,
associated_types,
extension,
}
}
}
impl<E: FnExtension + Default> From<&mut syn::TraitItemFn> for PublicFn<E> {
fn from(node: &mut syn::TraitItemFn) -> Self {
let payable = consume_flag(&mut node.attrs, "payable");
let selector_override =
consume_attr::<attrs::Selector>(&mut node.attrs, "selector").map(|s| s.value.value());
let fallback = consume_flag(&mut node.attrs, "fallback");
let receive = consume_flag(&mut node.attrs, "receive");
let constructor = consume_flag(&mut node.attrs, "constructor");
let kind = if fallback {
FnKind::Fallback {
with_args: node.sig.inputs.len() > 1,
}
} else if receive {
FnKind::Receive
} else if constructor {
FnKind::Constructor
} else {
FnKind::Function
};
let num_specials = (fallback as i8) + (constructor as i8) + (receive as i8);
if num_specials > 1 {
emit_error!(
node.span(),
"function can be only one of fallback, receive or constructor"
);
}
if num_specials > 0 && selector_override.is_some() {
emit_error!(
node.span(),
"fallback, receive, and constructor can't have custom selector"
);
}
let name = node.sig.ident.clone();
let (sol_name, name_err) = verify_sol_name(&kind, name.to_string(), selector_override);
if let Some(err) = name_err {
emit_error!(node.span(), err);
}
let sol_name = syn_solidity::SolIdent::new(&sol_name);
let (inferred_purity, has_self) = Purity::infer(&node.sig);
let purity = if payable || matches!(kind, FnKind::Receive) {
Purity::Payable
} else {
inferred_purity
};
let mut args = node.sig.inputs.iter();
if inferred_purity > Purity::Pure {
args.next();
}
let inputs = match kind {
FnKind::Function | FnKind::Constructor => args.map(PublicFnArg::from).collect(),
_ => Vec::new(),
};
let input_span = node.sig.inputs.span();
let output = match &node.sig.output {
syn::ReturnType::Default => None,
syn::ReturnType::Type(_, ty) => Some(*ty.clone()),
};
let output_span = output
.as_ref()
.map(Spanned::span)
.unwrap_or(node.sig.output.span());
let extension: E = E::default();
Self {
name,
sol_name,
purity,
inferred_purity,
kind,
has_self,
inputs,
input_span,
output: node.sig.output.clone(),
output_span,
extension,
}
}
}
impl<E: FnExtension> From<&mut syn::ImplItemFn> for PublicFn<E> {
fn from(node: &mut syn::ImplItemFn) -> Self {
let payable = consume_flag(&mut node.attrs, "payable");
let selector_override =
consume_attr::<attrs::Selector>(&mut node.attrs, "selector").map(|s| s.value.value());
let fallback = consume_flag(&mut node.attrs, "fallback");
let receive = consume_flag(&mut node.attrs, "receive");
let constructor = consume_flag(&mut node.attrs, "constructor");
let kind = if fallback {
FnKind::Fallback {
with_args: node.sig.inputs.len() > 1,
}
} else if receive {
FnKind::Receive
} else if constructor {
FnKind::Constructor
} else {
FnKind::Function
};
let num_specials = (fallback as i8) + (constructor as i8) + (receive as i8);
if num_specials > 1 {
emit_error!(
node.span(),
"function can be only one of fallback, receive or constructor"
);
}
if num_specials > 0 && selector_override.is_some() {
emit_error!(
node.span(),
"fallback, receive, and constructor can't have custom selector"
);
}
let name = node.sig.ident.clone();
let (sol_name, name_err) = verify_sol_name(&kind, name.to_string(), selector_override);
if let Some(err) = name_err {
emit_error!(node.span(), err);
}
let sol_name = syn_solidity::SolIdent::new(&sol_name);
let (inferred_purity, has_self) = Purity::infer(&node.sig);
let purity = if payable || matches!(kind, FnKind::Receive) {
Purity::Payable
} else {
inferred_purity
};
let mut args = node.sig.inputs.iter();
if inferred_purity > Purity::Pure {
args.next();
}
let inputs = match kind {
FnKind::Function | FnKind::Constructor => args.map(PublicFnArg::from).collect(),
_ => Vec::new(),
};
let input_span = node.sig.inputs.span();
let output = match &node.sig.output {
syn::ReturnType::Default => None,
syn::ReturnType::Type(_, ty) => Some(*ty.clone()),
};
let output_span = output
.as_ref()
.map(Spanned::span)
.unwrap_or(node.sig.output.span());
let extension = E::build(node);
Self {
name,
sol_name,
purity,
inferred_purity,
kind,
has_self,
inputs,
input_span,
output: node.sig.output.clone(),
output_span,
extension,
}
}
}
impl<E: FnArgExtension> From<&syn::FnArg> for PublicFnArg<E> {
fn from(node: &syn::FnArg) -> Self {
match node {
syn::FnArg::Typed(pat_type) => match &*pat_type.pat {
syn::Pat::Ident(pat_ident) => Self {
name: pat_ident.ident.clone(),
ty: *pat_type.ty.clone(),
extension: E::build(node),
},
other => {
emit_error!(other, "destructuring patterns are not supported in #[public] functions; use a named parameter instead");
Self {
name: syn::Ident::new("_", other.span()),
ty: parse_quote! { () },
extension: E::build(node),
}
}
},
syn::FnArg::Receiver(recv) => {
emit_error!(
recv,
"unexpected `self` parameter in #[public] function argument list"
);
Self {
name: syn::Ident::new("_", recv.span()),
ty: parse_quote! { () },
extension: E::build(node),
}
}
}
}
}
fn verify_sol_name(
kind: &FnKind,
name: String,
selector_override: Option<String>,
) -> (String, Option<String>) {
let name = selector_override.unwrap_or(name.to_case(Case::Camel));
let name_low = name.to_lowercase();
let err_kind = if name_low == "receive" && !matches!(kind, FnKind::Receive) {
Some("receive")
} else if name_low == "fallback" && !matches!(kind, FnKind::Fallback { .. }) {
Some("fallback")
} else if (name_low == "constructor" || name_low == "stylus_constructor")
&& !matches!(kind, FnKind::Constructor)
{
Some("constructor")
} else {
None
};
let err = err_kind.map(|kind_name| {
format!("{kind_name} function can only be defined using the corresponding attribute")
});
(name, err)
}
#[cfg(test)]
mod tests {
use quote::ToTokens;
use syn::parse_quote;
use super::{
types::{self, FnKind, PublicImpl, PublicTrait},
verify_sol_name,
};
#[test]
fn test_public_consumes_payable() {
let mut impl_item = parse_quote! {
#[derive(Debug)]
impl Contract {
#[payable]
#[other]
fn func() {}
}
};
let _public = PublicImpl::from(&mut impl_item);
let syn::ImplItem::Fn(syn::ImplItemFn { attrs, .. }) = &impl_item.items[0] else {
unreachable!();
};
assert_eq!(attrs, &vec![parse_quote! { #[other] }]);
}
#[test]
fn test_public_consumes_constructor() {
let mut impl_item = parse_quote! {
#[derive(Debug)]
impl Contract {
#[constructor]
fn func(&mut self, val: U256) {}
}
};
let public = PublicImpl::from(&mut impl_item);
assert!(matches!(public.funcs[0].kind, FnKind::Constructor));
let syn::ImplItem::Fn(syn::ImplItemFn { attrs, .. }) = &impl_item.items[0] else {
unreachable!();
};
assert!(attrs.is_empty());
}
#[test]
fn test_verify_sol_name() {
let cases = vec![
("foo", None, "foo", false),
("foo_bar", None, "fooBar", false),
("foo_baz", Some("fooBar"), "fooBar", false),
("foo_baz", Some("fooBAR"), "fooBAR", false),
("receive", None, "receive", true),
("re_ceive", None, "reCeive", true),
("foo", Some("RECEIVE"), "RECEIVE", true),
];
for (name, selector_override, expected_sol_name, has_err) in cases {
let kind = FnKind::Function;
let (sol_name, err) =
verify_sol_name(&kind, name.to_owned(), selector_override.map(String::from));
assert_eq!(sol_name, expected_sol_name);
assert_eq!(err.is_some(), has_err);
}
}
#[test]
fn test_display_label() {
let mut impl_item: syn::ItemImpl = parse_quote! {
impl Contract {
fn foo(_x: u64) {}
}
};
let public = PublicImpl::from(&mut impl_item);
assert_eq!(public.funcs[0].display_label(), "`foo`");
let mut impl_item: syn::ItemImpl = parse_quote! {
impl Contract {
fn foo_bar(_x: u64) {}
}
};
let public = PublicImpl::from(&mut impl_item);
assert_eq!(
public.funcs[0].display_label(),
"`foo_bar` (ABI name `fooBar`)"
);
}
#[test]
fn test_selector_collision_checks_count() {
let checks = types::selector_collision_checks::<()>(&[]);
assert!(checks.is_empty(), "expected 0 checks with empty input");
let mut impl_item: syn::ItemImpl = parse_quote! {
impl Contract {
fn solo(_x: u64) {}
}
};
let public = PublicImpl::from(&mut impl_item);
let checks = types::selector_collision_checks(&public.funcs);
assert!(
checks.is_empty(),
"expected 0 checks with a single function"
);
let mut impl_item: syn::ItemImpl = parse_quote! {
impl Contract {
fn foo_bar(_x: u64) {}
#[allow(non_snake_case)]
fn fooBar(_x: u64) {}
}
};
let public = PublicImpl::from(&mut impl_item);
let checks = types::selector_collision_checks(&public.funcs);
assert_eq!(checks.len(), 1, "expected 1 pairwise collision check");
let mut impl_item: syn::ItemImpl = parse_quote! {
impl Contract {
fn alpha(_x: u64) {}
fn beta(_x: u64) {}
fn gamma(_x: u64) {}
}
};
let public = PublicImpl::from(&mut impl_item);
let checks = types::selector_collision_checks(&public.funcs);
assert_eq!(checks.len(), 3, "expected 3 pairwise collision checks");
let mut impl_item: syn::ItemImpl = parse_quote! {
impl Contract {
fn alpha(_x: u64) {}
fn beta(_x: u64) {}
fn gamma(_x: u64) {}
fn delta(_x: u64) {}
}
};
let public = PublicImpl::from(&mut impl_item);
let checks = types::selector_collision_checks(&public.funcs);
assert_eq!(checks.len(), 6, "expected 6 pairwise collision checks");
let mut impl_item: syn::ItemImpl = parse_quote! {
impl Contract {
fn alpha(_x: u64) {}
#[fallback]
fn my_fallback(&mut self, _args: &[u8]) -> stylus_sdk::ArbResult { Ok(vec![]) }
#[receive]
fn my_receive(&mut self) -> Result<(), Vec<u8>> { Ok(()) }
}
};
let public = PublicImpl::from(&mut impl_item);
let checks = types::selector_collision_checks(&public.funcs);
assert!(
checks.is_empty(),
"expected 0 checks with only one regular function"
);
let mut impl_item: syn::ItemImpl = parse_quote! {
impl Contract {
fn alpha(_x: u64) {}
fn beta(_x: u64) {}
#[constructor]
fn my_constructor(&mut self) {}
}
};
let public = PublicImpl::from(&mut impl_item);
let checks = types::selector_collision_checks(&public.funcs);
assert_eq!(
checks.len(),
1,
"expected 1 check: constructor excluded, 2 regular functions remain"
);
}
#[test]
fn test_selector_collision_checks_trait() {
let mut trait_item: syn::ItemTrait = parse_quote! {
trait MyContract {
fn foo(_x: u64) {}
fn bar(_x: u64) {}
fn baz(_x: u64) {}
}
};
let public = PublicTrait::from(&mut trait_item);
let checks = types::selector_collision_checks(&public.funcs);
assert_eq!(
checks.len(),
3,
"expected 3 pairwise collision checks for trait"
);
let mut trait_item: syn::ItemTrait = parse_quote! {
trait MyContract {
fn only_one(_x: u64) {}
}
};
let public = PublicTrait::from(&mut trait_item);
let checks = types::selector_collision_checks(&public.funcs);
assert!(
checks.is_empty(),
"expected 0 checks for single-method trait"
);
}
#[test]
fn test_selector_collision_checks_content() {
let mut impl_item: syn::ItemImpl = parse_quote! {
impl Contract {
fn foo(_x: u64) {}
fn bar(_x: u64) {}
}
};
let public = PublicImpl::from(&mut impl_item);
let checks = types::selector_collision_checks(&public.funcs);
assert_eq!(checks.len(), 1);
let tokens = checks[0].to_token_stream().to_string();
assert!(
tokens.contains("__SELECTOR_foo"),
"expected __SELECTOR_foo in generated check"
);
assert!(
tokens.contains("__SELECTOR_bar"),
"expected __SELECTOR_bar in generated check"
);
assert!(
tokens.contains("contract-client-gen"),
"expected cfg gate for contract-client-gen"
);
assert!(
tokens.contains("ABI selector collision"),
"expected collision error message in generated check"
);
}
#[test]
fn test_selector_collision_checks_content_camel_case() {
let mut impl_item: syn::ItemImpl = parse_quote! {
impl Contract {
fn foo_bar(_x: u64) {}
fn baz_qux(_x: u64) {}
}
};
let public = PublicImpl::from(&mut impl_item);
let checks = types::selector_collision_checks(&public.funcs);
assert_eq!(checks.len(), 1);
let tokens = checks[0].to_token_stream().to_string();
assert!(
tokens.contains("foo_bar"),
"expected Rust name foo_bar in error message"
);
assert!(
tokens.contains("fooBar"),
"expected ABI name fooBar in error message"
);
assert!(
tokens.contains("baz_qux"),
"expected Rust name baz_qux in error message"
);
assert!(
tokens.contains("bazQux"),
"expected ABI name bazQux in error message"
);
}
#[test]
fn test_selector_collision_checks_content_with_selector_override() {
let mut impl_item: syn::ItemImpl = parse_quote! {
impl Contract {
#[selector(name = "customName")]
fn my_func(_x: u64) {}
fn other(_x: u64) {}
}
};
let public = PublicImpl::from(&mut impl_item);
let checks = types::selector_collision_checks(&public.funcs);
assert_eq!(checks.len(), 1);
let tokens = checks[0].to_token_stream().to_string();
assert!(
tokens.contains("__SELECTOR_my_func"),
"expected __SELECTOR_my_func in generated check"
);
assert!(
tokens.contains("customName"),
"expected overridden ABI name customName in error message"
);
}
#[test]
fn test_selector_collision_checks_all_pairs_covered() {
let mut impl_item: syn::ItemImpl = parse_quote! {
impl Contract {
fn alpha(_x: u64) {}
fn beta(_x: u64) {}
fn gamma(_x: u64) {}
}
};
let public = PublicImpl::from(&mut impl_item);
let checks = types::selector_collision_checks(&public.funcs);
assert_eq!(checks.len(), 3);
let tokens: Vec<String> = checks
.iter()
.map(|c| c.to_token_stream().to_string())
.collect();
let has_pair = |a: &str, b: &str| {
tokens.iter().any(|t| {
t.contains(&format!("__SELECTOR_{a}")) && t.contains(&format!("__SELECTOR_{b}"))
})
};
assert!(has_pair("alpha", "beta"), "missing alpha-beta pair");
assert!(has_pair("alpha", "gamma"), "missing alpha-gamma pair");
assert!(has_pair("beta", "gamma"), "missing beta-gamma pair");
}
}