use proc_macro::TokenStream;
use quote::quote;
#[proc_macro_derive(MultiSenderFrom)]
pub fn derive_multi_sender_from(input: TokenStream) -> TokenStream {
derive_multi_sender_from_impl(input.into()).into()
}
fn derive_multi_sender_from_impl(input: proc_macro2::TokenStream) -> proc_macro2::TokenStream {
let ast: syn::DeriveInput = syn::parse2(input).unwrap();
let struct_name = ast.ident.clone();
let input = match ast.data {
syn::Data::Struct(input) => input,
_ => {
panic!("MultiSenderFrom can only be derived for structs");
}
};
let mut type_bounds = Vec::new();
let mut initializers = Vec::new();
let mut cfg_attrs = Vec::new();
let mut names = Vec::<syn::Ident>::new();
for (i, field) in input.fields.into_iter().enumerate() {
let field_name = field
.ident
.as_ref()
.map(|ident| ident.to_string())
.unwrap_or_else(|| format!("#{}", i));
cfg_attrs.push(extract_cfg_attributes(&field.attrs));
match &field.ty {
syn::Type::Path(path) => {
let last_segment = path.path.segments.last().unwrap();
let arguments = match last_segment.arguments.clone() {
syn::PathArguments::AngleBracketed(arguments) => {
arguments.args.into_iter().collect::<Vec<_>>()
}
_ => panic!("Field {} must be either a Sender or an AsyncSender", field_name),
};
if last_segment.ident == "Sender" {
type_bounds.push(quote!(near_async::messaging::CanSend<#(#arguments),*>));
initializers.push(quote!(near_async::messaging::IntoSender::as_sender(&input)));
} else if last_segment.ident == "AsyncSender" {
type_bounds.push(quote!(
near_async::messaging::CanSendAsync<#(#arguments),*>
));
initializers.push(quote!(
near_async::messaging::IntoAsyncSender::as_async_sender(&input)
));
} else {
panic!("Field {} must be either a Sender or an AsyncSender", field_name);
}
if let Some(name) = &field.ident {
names.push(name.clone());
}
}
_ => panic!("Field {} must be either a Sender or an AsyncSender", field_name),
}
}
assert!(!type_bounds.is_empty(), "Must have at least one field");
let initializer = if names.is_empty() {
quote!(#struct_name(#(#(#cfg_attrs)* #initializers,)*))
} else {
quote!(#struct_name {
#(#(#cfg_attrs)* #names: #initializers,)*
})
};
quote! {
impl<A: #(#type_bounds)+*> near_async::messaging::MultiSenderFrom<A> for #struct_name {
fn multi_sender_from(input: std::sync::Arc<A>) -> Self {
#initializer
}
}
}
}
#[proc_macro_derive(MultiSend)]
pub fn derive_multi_send(input: TokenStream) -> TokenStream {
derive_multi_send_impl(input.into()).into()
}
fn derive_multi_send_impl(input: proc_macro2::TokenStream) -> proc_macro2::TokenStream {
let ast: syn::DeriveInput = syn::parse2(input).unwrap();
let struct_name = ast.ident.clone();
let input = match ast.data {
syn::Data::Struct(input) => input,
_ => {
panic!("MultiSend can only be derived for structs");
}
};
let mut tokens = Vec::new();
for (i, field) in input.fields.into_iter().enumerate() {
let field_name = field.ident.as_ref().map(|ident| quote!(#ident)).unwrap_or_else(|| {
let index = syn::Index::from(i);
quote!(#index)
});
let cfg_attrs = extract_cfg_attributes(&field.attrs);
if let syn::Type::Path(path) = &field.ty {
let last_segment = path.path.segments.last().unwrap();
let arguments = match last_segment.arguments.clone() {
syn::PathArguments::AngleBracketed(arguments) => {
arguments.args.into_iter().collect::<Vec<_>>()
}
_ => {
continue;
}
};
if last_segment.ident == "Sender" {
let message_type = arguments[0].clone();
tokens.push(quote! {
#(#cfg_attrs)*
impl near_async::messaging::CanSend<#message_type> for #struct_name {
fn send(&self, message: #message_type) {
self.#field_name.send(message);
}
}
});
} else if last_segment.ident == "AsyncSender" {
let message_type = arguments[0].clone();
let result_type = arguments[1].clone();
tokens.push(quote! {
#(#cfg_attrs)*
impl near_async::messaging::CanSendAsync<#message_type, #result_type> for #struct_name {
fn send_async(&self, message: #message_type)
-> near_async::futures::BoxFuture<'static, Result<#result_type, near_async::messaging::AsyncSendError>>
{
self.#field_name.send_async(message)
}
}
});
}
}
}
quote! {#(#tokens)*}
}
fn extract_cfg_attributes(attrs: &[syn::Attribute]) -> Vec<syn::Attribute> {
attrs.iter().filter(|attr| attr.path().is_ident("cfg")).cloned().collect()
}
#[cfg(test)]
mod tests {
use quote::quote;
#[test]
fn test_derive_into_multi_send() {
let input = quote! {
struct TestSenders {
sender: Sender<String>,
async_sender: AsyncSender<String, u32>,
qualified_sender: near_async::messaging::Sender<i32>,
qualified_async_sender: near_async::messaging::AsyncSender<i32, String>,
}
};
let expected = quote! {
impl<A:
near_async::messaging::CanSend<String>
+ near_async::messaging::CanSendAsync<String, u32>
+ near_async::messaging::CanSend<i32>
+ near_async::messaging::CanSendAsync<i32, String>
> near_async::messaging::MultiSenderFrom<A> for TestSenders {
fn multi_sender_from(input: std::sync::Arc<A>) -> Self {
TestSenders {
sender: near_async::messaging::IntoSender::as_sender(&input),
async_sender: near_async::messaging::IntoAsyncSender::as_async_sender(&input),
qualified_sender: near_async::messaging::IntoSender::as_sender(&input),
qualified_async_sender: near_async::messaging::IntoAsyncSender::as_async_sender(&input),
}
}
}
};
let actual = super::derive_multi_sender_from_impl(input);
pretty_assertions::assert_str_eq!(actual.to_string(), expected.to_string());
}
#[test]
fn test_derive_multi_send() {
let input = quote! {
struct TestSenders {
sender: Sender<String>,
async_sender: AsyncSender<String, u32>,
qualified_sender: near_async::messaging::Sender<i32>,
qualified_async_sender: near_async::messaging::AsyncSender<i32, String>,
}
};
let expected = quote! {
impl near_async::messaging::CanSend<String> for TestSenders {
fn send(&self, message: String) {
self.sender.send(message);
}
}
impl near_async::messaging::CanSendAsync<String, u32> for TestSenders {
fn send_async(&self, message: String) -> near_async::futures::BoxFuture<'static, Result<u32, near_async::messaging::AsyncSendError>> {
self.async_sender.send_async(message)
}
}
impl near_async::messaging::CanSend<i32> for TestSenders {
fn send(&self, message: i32) {
self.qualified_sender.send(message);
}
}
impl near_async::messaging::CanSendAsync<i32, String> for TestSenders {
fn send_async(&self, message: i32) -> near_async::futures::BoxFuture<'static, Result<String, near_async::messaging::AsyncSendError>> {
self.qualified_async_sender.send_async(message)
}
}
};
let actual = super::derive_multi_send_impl(input);
pretty_assertions::assert_str_eq!(actual.to_string(), expected.to_string());
}
}