extern crate proc_macro;
use {
anchor_syn::{codegen::program::common::gen_discriminator, Overrides},
quote::{quote, ToTokens},
syn::{
parenthesized,
parse::{Parse, ParseStream},
parse_macro_input,
token::Paren,
Ident, LitStr, Token,
},
};
mod id;
#[cfg(feature = "lazy-account")]
mod lazy;
#[proc_macro_attribute]
pub fn account(
args: proc_macro::TokenStream,
input: proc_macro::TokenStream,
) -> proc_macro::TokenStream {
let args = parse_macro_input!(args as AccountArgs);
let namespace = args.namespace.unwrap_or_default();
let is_zero_copy = args.zero_copy.is_some();
let unsafe_bytemuck = args.zero_copy.unwrap_or_default();
let account_strct = parse_macro_input!(input as syn::ItemStruct);
let account_name = &account_strct.ident;
let account_name_str = account_name.to_string();
let (impl_gen, type_gen, where_clause) = account_strct.generics.split_for_impl();
let discriminator = args
.overrides
.and_then(|ov| ov.discriminator)
.unwrap_or_else(|| {
let namespace = if namespace.is_empty() {
"account"
} else {
&namespace
};
gen_discriminator(namespace, account_name)
});
let disc = if account_strct.generics.lt_token.is_some() {
quote! { #account_name::#type_gen::DISCRIMINATOR }
} else {
quote! { #account_name::DISCRIMINATOR }
};
let owner_impl = {
if namespace.is_empty() {
quote! {
#[automatically_derived]
impl #impl_gen anchor_lang::Owner for #account_name #type_gen #where_clause {
fn owner() -> Pubkey {
#[cfg(not(doctest))]
{ crate::ID }
#[cfg(doctest)]
{ ID }
}
}
}
} else {
quote! {}
}
};
let unsafe_bytemuck_impl = {
if unsafe_bytemuck {
quote! {
#[automatically_derived]
unsafe impl #impl_gen anchor_lang::__private::bytemuck::Pod for #account_name #type_gen #where_clause {}
#[automatically_derived]
unsafe impl #impl_gen anchor_lang::__private::bytemuck::Zeroable for #account_name #type_gen #where_clause {}
}
} else {
quote! {}
}
};
let bytemuck_derives = {
if !unsafe_bytemuck {
quote! {
#[zero_copy]
}
} else {
quote! {
#[zero_copy(unsafe)]
}
}
};
proc_macro::TokenStream::from({
if is_zero_copy {
quote! {
#bytemuck_derives
#account_strct
#unsafe_bytemuck_impl
#[automatically_derived]
impl #impl_gen anchor_lang::ZeroCopy for #account_name #type_gen #where_clause {}
#[automatically_derived]
impl #impl_gen anchor_lang::Discriminator for #account_name #type_gen #where_clause {
const DISCRIMINATOR: &'static [u8] = #discriminator;
}
#[automatically_derived]
impl #impl_gen anchor_lang::AccountDeserialize for #account_name #type_gen #where_clause {
fn try_deserialize(buf: &mut &[u8]) -> anchor_lang::Result<Self> {
if buf.len() < #disc.len() {
return Err(anchor_lang::error::ErrorCode::AccountDiscriminatorNotFound.into());
}
let given_disc = &buf[..#disc.len()];
if #disc != given_disc {
return Err(anchor_lang::error!(anchor_lang::error::ErrorCode::AccountDiscriminatorMismatch).with_account_name(#account_name_str));
}
Self::try_deserialize_unchecked(buf)
}
fn try_deserialize_unchecked(buf: &mut &[u8]) -> anchor_lang::Result<Self> {
let data: &[u8] = &buf[#disc.len()..];
Ok(anchor_lang::__private::bytemuck::pod_read_unaligned(data))
}
}
#owner_impl
}
} else {
let lazy = {
#[cfg(feature = "lazy-account")]
match namespace.is_empty().then(|| lazy::gen_lazy(&account_strct)) {
Some(Ok(lazy)) => lazy,
_ => Default::default(),
}
#[cfg(not(feature = "lazy-account"))]
proc_macro2::TokenStream::default()
};
quote! {
#[derive(AnchorSerialize, AnchorDeserialize, Clone)]
#account_strct
#[automatically_derived]
impl #impl_gen anchor_lang::AccountSerialize for #account_name #type_gen #where_clause {
fn try_serialize<W: std::io::Write>(&self, writer: &mut W) -> anchor_lang::Result<()> {
if writer.write_all(#disc).is_err() {
return Err(anchor_lang::error::ErrorCode::AccountDidNotSerialize.into());
}
if AnchorSerialize::serialize(self, writer).is_err() {
return Err(anchor_lang::error::ErrorCode::AccountDidNotSerialize.into());
}
Ok(())
}
}
#[automatically_derived]
impl #impl_gen anchor_lang::AccountDeserialize for #account_name #type_gen #where_clause {
fn try_deserialize(buf: &mut &[u8]) -> anchor_lang::Result<Self> {
if buf.len() < #disc.len() {
return Err(anchor_lang::error::ErrorCode::AccountDiscriminatorNotFound.into());
}
let given_disc = &buf[..#disc.len()];
if #disc != given_disc {
return Err(anchor_lang::error!(anchor_lang::error::ErrorCode::AccountDiscriminatorMismatch).with_account_name(#account_name_str));
}
Self::try_deserialize_unchecked(buf)
}
fn try_deserialize_unchecked(buf: &mut &[u8]) -> anchor_lang::Result<Self> {
let mut data: &[u8] = &buf[#disc.len()..];
AnchorDeserialize::deserialize(&mut data)
.map_err(|_| anchor_lang::error::ErrorCode::AccountDidNotDeserialize.into())
}
}
#[automatically_derived]
impl #impl_gen anchor_lang::Discriminator for #account_name #type_gen #where_clause {
const DISCRIMINATOR: &'static [u8] = #discriminator;
}
#owner_impl
#lazy
}
}
})
}
#[derive(Debug, Default)]
struct AccountArgs {
zero_copy: Option<bool>,
namespace: Option<String>,
overrides: Option<Overrides>,
}
impl Parse for AccountArgs {
fn parse(input: ParseStream) -> syn::Result<Self> {
let mut parsed = Self::default();
let args = input.parse_terminated(AccountArg::parse, Token![,])?;
for arg in args {
match arg {
AccountArg::ZeroCopy { is_unsafe } => {
parsed.zero_copy.replace(is_unsafe);
}
AccountArg::Namespace(ns) => {
parsed.namespace.replace(ns);
}
AccountArg::Overrides(ov) => {
parsed.overrides.replace(ov);
}
}
}
Ok(parsed)
}
}
enum AccountArg {
ZeroCopy { is_unsafe: bool },
Namespace(String),
Overrides(Overrides),
}
impl Parse for AccountArg {
fn parse(input: ParseStream) -> syn::Result<Self> {
if let Ok(ns) = input.parse::<LitStr>() {
return Ok(Self::Namespace(
ns.to_token_stream().to_string().replace('\"', ""),
));
}
if input
.fork()
.parse::<Ident>()
.is_ok_and(|ident| ident == "zero_copy")
{
input.parse::<Ident>()?;
let is_unsafe = if input.peek(Paren) {
let content;
parenthesized!(content in input);
let content = content.parse::<proc_macro2::TokenStream>()?;
if content.to_string().as_str().trim() != "unsafe" {
return Err(syn::Error::new(
syn::spanned::Spanned::span(&content),
"Expected `unsafe`",
));
}
true
} else {
false
};
return Ok(Self::ZeroCopy { is_unsafe });
}
input.parse::<Overrides>().map(Self::Overrides)
}
}
#[proc_macro_derive(ZeroCopyAccessor, attributes(accessor))]
pub fn derive_zero_copy_accessor(item: proc_macro::TokenStream) -> proc_macro::TokenStream {
let account_strct = parse_macro_input!(item as syn::ItemStruct);
let account_name = &account_strct.ident;
let (impl_gen, ty_gen, where_clause) = account_strct.generics.split_for_impl();
let fields = match &account_strct.fields {
syn::Fields::Named(n) => n,
_ => {
return syn::Error::new_spanned(
&account_strct.ident,
"#[derive(ZeroCopyAccessor)] requires a struct with named fields",
)
.into_compile_error()
.into()
}
};
let methods: Vec<proc_macro2::TokenStream> = fields
.named
.iter()
.filter_map(|field: &syn::Field| {
field
.attrs
.iter()
.find(|attr| anchor_syn::parser::tts_to_string(attr.path()) == "accessor")
.map(|attr| {
let tokens = match &attr.meta {
syn::Meta::List(list) => list.tokens.clone(),
_ => {
return syn::Error::new_spanned(
attr,
"`#[accessor]` requires a type argument, e.g `#[accessor(MyType)]`",
)
.into_compile_error();
}
};
let accessor_ty = match tokens.into_iter().next() {
Some(token) => token,
None => {
return syn::Error::new_spanned(
attr,
"`#[accessor]` requires a type inside the parentheses e.g \
`#[accessor(MyType)]`",
)
.into_compile_error()
}
};
#[allow(
clippy::unwrap_used,
reason = "accessor fields always have idents (named struct fields)"
)]
let field_name = field.ident.as_ref().unwrap();
#[allow(
clippy::unwrap_used,
reason = "get_<field_name> formed from a valid Rust identifier is always \
valid TokenStream"
)]
let get_field: proc_macro2::TokenStream =
format!("get_{field_name}").parse().unwrap();
#[allow(
clippy::unwrap_used,
reason = "set_<field_name> formed from a valid Rust identifier is always \
valid TokenStream"
)]
let set_field: proc_macro2::TokenStream =
format!("set_{field_name}").parse().unwrap();
quote! {
pub fn #get_field(&self) -> #accessor_ty {
anchor_lang::__private::ZeroCopyAccessor::get(&self.#field_name)
}
pub fn #set_field(&mut self, input: &#accessor_ty) {
self.#field_name = anchor_lang::__private::ZeroCopyAccessor::set(input);
}
}
})
})
.collect();
proc_macro::TokenStream::from(quote! {
#[automatically_derived]
impl #impl_gen #account_name #ty_gen #where_clause {
#(#methods)*
}
})
}
#[proc_macro_attribute]
pub fn zero_copy(
args: proc_macro::TokenStream,
item: proc_macro::TokenStream,
) -> proc_macro::TokenStream {
let mut is_unsafe = false;
for arg in args.into_iter() {
match arg {
proc_macro::TokenTree::Ident(ident) => {
if ident.to_string() == "unsafe" {
is_unsafe = true;
} else {
return syn::Error::new(
proc_macro2::Span::from(ident.span()),
"expected `unsafe`, e.g `#[zero_copy(unsafe)]`",
)
.into_compile_error()
.into();
}
}
_ => {
return syn::Error::new(
proc_macro2::Span::from(arg.span()),
"expected `unsafe`, e.g `#[zero_copy(unsafe)]`",
)
.into_compile_error()
.into();
}
}
}
let account_strct = parse_macro_input!(item as syn::ItemStruct);
let attr = account_strct
.attrs
.iter()
.find(|attr| anchor_syn::parser::tts_to_string(attr.path()) == "repr");
let repr = match attr {
Some(_attr) => quote! {},
None => {
if is_unsafe {
quote! {#[repr(Rust, packed)]}
} else {
quote! {#[repr(C)]}
}
}
};
let mut has_pod_attr = false;
let mut has_zeroable_attr = false;
for attr in account_strct.attrs.iter() {
if !attr.path().is_ident("derive") {
continue;
}
if let syn::Meta::List(list) = &attr.meta {
let tokens_str = list.tokens.to_string();
if tokens_str.contains("bytemuck :: Pod") {
has_pod_attr = true;
}
if tokens_str.contains("bytemuck :: Zeroable") {
has_zeroable_attr = true;
}
}
}
let pod = if has_pod_attr || is_unsafe {
quote! {}
} else {
quote! {#[derive(::bytemuck::Pod)]}
};
let zeroable = if has_zeroable_attr || is_unsafe {
quote! {}
} else {
quote! {#[derive(::bytemuck::Zeroable)]}
};
let ret = quote! {
#[derive(anchor_lang::__private::ZeroCopyAccessor, Copy, Clone)]
#repr
#pod
#zeroable
#account_strct
};
#[cfg(feature = "idl-build")]
{
let derive_unsafe = if is_unsafe {
quote! { #[derive(bytemuck::Unsafe)] }
} else {
quote! {}
};
let zc_struct = syn::parse_quote! {
#derive_unsafe
#ret
};
let idl_build_impl = anchor_syn::idl::impl_idl_build_struct(&zc_struct);
return proc_macro::TokenStream::from(quote! {
#ret
#idl_build_impl
});
}
#[allow(unreachable_code)]
proc_macro::TokenStream::from(ret)
}
#[proc_macro]
pub fn pubkey(input: proc_macro::TokenStream) -> proc_macro::TokenStream {
let pk = parse_macro_input!(input as id::Pubkey);
proc_macro::TokenStream::from(quote! {#pk})
}
#[proc_macro]
pub fn declare_id(input: proc_macro::TokenStream) -> proc_macro::TokenStream {
#[cfg(feature = "idl-build")]
let address = input.clone().to_string();
let id = parse_macro_input!(input as id::Id);
let ret = quote! { #id };
#[cfg(feature = "idl-build")]
{
let idl_print = anchor_syn::idl::gen_idl_print_fn_address(address);
return proc_macro::TokenStream::from(quote! {
#ret
#idl_print
});
}
#[allow(unreachable_code)]
proc_macro::TokenStream::from(ret)
}