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::{Comma, Paren},
Ident, LitStr,
},
};
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 {
crate::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()..];
let account = anchor_lang::__private::bytemuck::from_bytes(data);
Ok(*account)
}
}
#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::<_, Comma>(AccountArg::parse)?;
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>()? == "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,
_ => panic!("Fields must be named"),
};
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 mut tts = attr.tokens.clone().into_iter();
let g_stream = match tts.next().expect("Must have a token group") {
proc_macro2::TokenTree::Group(g) => g.stream(),
_ => panic!("Invalid syntax"),
};
let accessor_ty = match g_stream.into_iter().next() {
Some(token) => token,
_ => panic!("Missing accessor type"),
};
let field_name = field.ident.as_ref().unwrap();
let get_field: proc_macro2::TokenStream =
format!("get_{field_name}").parse().unwrap();
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 {
panic!("expected single ident `unsafe`");
}
}
_ => {
panic!("expected single ident `unsafe`");
}
}
}
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() {
let token_string = attr.tokens.to_string();
if token_string.contains("bytemuck :: Pod") {
has_pod_attr = true;
}
if token_string.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::parse2(quote! {
#derive_unsafe
#ret
})
.unwrap();
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)
}