#![feature(proc_macro_def_site)]
extern crate proc_macro;
use convert_case::Case;
use convert_case::Casing;
use indoc::indoc;
use proc_macro::TokenStream;
use proc_macro2::Span;
use quote::format_ident;
use quote::quote;
use syn::Attribute;
use syn::Data;
use syn::DataEnum;
use syn::DeriveInput;
use syn::Expr;
use syn::Field;
use syn::Fields;
use syn::Ident;
use syn::ItemFn;
use syn::ItemImpl;
use syn::Lit;
use syn::Meta;
use syn::MetaNameValue;
use syn::Token;
use syn::Type;
use syn::parse_macro_input;
use syn::punctuated::Punctuated;
use syn::spanned::Spanned;
const REPLY_VARIANT_ERROR: &str = indoc! {r#"
`call` message expects a typed `OncePortRef` or `OncePortHandle` argument in the last position
= help: use `MyCall(Arg1Type, Arg2Type, .., OncePortRef<ReplyType>)`
= help: use `MyCall(Arg1Type, Arg2Type, .., OncePortHandle<ReplyType>)`
"#};
const REPLY_USAGE_ERROR: &str = indoc! {r#"
`call` message expects at most one `reply` argument
= help: use `MyCall(Arg1Type, Arg2Type, .., #[reply] OncePortRef<ReplyType>)`
= help: use `MyCall(Arg1Type, Arg2Type, .., #[reply] OncePortHandle<ReplyType>)`
"#};
enum FieldFlag {
None,
Reply,
}
enum Variant {
Named {
enum_name: Ident,
name: Ident,
field_names: Vec<Ident>,
field_types: Vec<Type>,
field_flags: Vec<FieldFlag>,
},
Anon {
enum_name: Ident,
name: Ident,
field_types: Vec<Type>,
field_flags: Vec<FieldFlag>,
},
}
impl Variant {
fn len(&self) -> usize {
self.field_types().len()
}
fn enum_name(&self) -> &Ident {
match self {
Variant::Named { enum_name, .. } => enum_name,
Variant::Anon { enum_name, .. } => enum_name,
}
}
fn name(&self) -> &Ident {
match self {
Variant::Named { name, .. } => name,
Variant::Anon { name, .. } => name,
}
}
fn snake_name(&self) -> Ident {
Ident::new(
&self.name().to_string().to_case(Case::Snake),
self.name().span(),
)
}
fn qualified_name(&self) -> proc_macro2::TokenStream {
let enum_name = self.enum_name();
let name = self.name();
quote! { #enum_name::#name }
}
fn field_names(&self) -> Vec<Ident> {
match self {
Variant::Named { field_names, .. } => field_names.clone(),
Variant::Anon { field_types, .. } => (0usize..field_types.len())
.map(|idx| format_ident!("arg{}", idx))
.collect(),
}
}
fn field_types(&self) -> &Vec<Type> {
match self {
Variant::Named { field_types, .. } => field_types,
Variant::Anon { field_types, .. } => field_types,
}
}
fn field_flags(&self) -> &Vec<FieldFlag> {
match self {
Variant::Named { field_flags, .. } => field_flags,
Variant::Anon { field_flags, .. } => field_flags,
}
}
fn constructor(&self) -> proc_macro2::TokenStream {
let qualified_name = self.qualified_name();
let field_names = self.field_names();
match self {
Variant::Named { .. } => quote! { #qualified_name { #(#field_names),* } },
Variant::Anon { .. } => quote! { #qualified_name(#(#field_names),*) },
}
}
}
enum Message {
Call {
variant: Variant,
reply_port_is_handle: bool,
return_type: Type,
log_level: Option<Ident>,
},
OneWay {
variant: Variant,
log_level: Option<Ident>,
},
}
impl Message {
fn new(span: Span, variant: Variant, log_level: Option<Ident>) -> Result<Self, syn::Error> {
match &variant
.field_flags()
.iter()
.zip(variant.field_types())
.filter_map(|(flag, ty)| match flag {
FieldFlag::Reply => Some(ty),
FieldFlag::None => None,
})
.collect::<Vec<&Type>>()[..]
{
[] => Ok(Self::OneWay { variant, log_level }),
[reply_port_ty] => {
let syn::Type::Path(type_path) = reply_port_ty else {
return Err(syn::Error::new(span, REPLY_VARIANT_ERROR));
};
let Some(last_segment) = type_path.path.segments.last() else {
return Err(syn::Error::new(span, REPLY_VARIANT_ERROR));
};
if last_segment.ident != "OncePortRef" && last_segment.ident != "OncePortHandle" {
return Err(syn::Error::new_spanned(last_segment, REPLY_VARIANT_ERROR));
}
let syn::PathArguments::AngleBracketed(args) = &last_segment.arguments else {
return Err(syn::Error::new_spanned(last_segment, REPLY_VARIANT_ERROR));
};
let Some(syn::GenericArgument::Type(return_ty)) = args.args.first() else {
return Err(syn::Error::new_spanned(&args.args, REPLY_VARIANT_ERROR));
};
let reply_port_is_handle = last_segment.ident == "OncePortHandle";
let return_type = return_ty.clone();
Ok(Self::Call {
variant,
reply_port_is_handle,
return_type,
log_level,
})
}
_ => Err(syn::Error::new(span, REPLY_USAGE_ERROR)),
}
}
fn args(&self) -> Vec<(Ident, Type)> {
match self {
Message::Call { variant, .. } => variant
.field_names()
.into_iter()
.zip(variant.field_types().clone())
.take(variant.len() - 1)
.collect(),
Message::OneWay { variant, .. } => variant
.field_names()
.into_iter()
.zip(variant.field_types().clone())
.collect(),
}
}
fn variant(&self) -> &Variant {
match self {
Message::Call { variant, .. } => variant,
Message::OneWay { variant, .. } => variant,
}
}
fn reply_port_position(&self) -> Option<usize> {
self.variant()
.field_flags()
.iter()
.position(|flag| matches!(flag, FieldFlag::Reply))
}
fn reply_port_arg(&self) -> Option<(Ident, Type)> {
match self {
Message::Call { variant, .. } => {
let pos = self.reply_port_position()?;
Some((
variant.field_names()[pos].clone(),
variant.field_types()[pos].clone(),
))
}
Message::OneWay { .. } => None,
}
}
}
fn parse_log_level(attrs: &[Attribute]) -> Result<Option<Ident>, syn::Error> {
let level: Option<String> = match attrs.iter().find(|attr| attr.path().is_ident("log_level")) {
Some(attr) => {
let Ok(meta) = attr.meta.require_list() else {
return Err(syn::Error::new(
Span::call_site(),
indoc! {"
`log_level` attribute must specify level. Supported levels = error, warn, info, debug, trace
= help use `#[log_level(info)]` or `#[log_level(error)]`
"},
));
};
let parsed = meta.parse_args_with(Punctuated::<Ident, Token![,]>::parse_terminated)?;
if parsed.len() != 1 {
return Err(syn::Error::new(
Span::call_site(),
indoc! {"
`log_level` attribute must specify exactly one level
= help use `#[log_level(warn)]` or `#[log_level(info)]`
"},
));
};
Some(parsed.first().unwrap().to_string())
}
None => None,
};
if level.is_none() {
return Ok(None);
}
let level = level.unwrap();
match level.as_str() {
"error" | "warn" | "info" | "debug" | "trace" => {}
_ => {
return Err(syn::Error::new(
Span::call_site(),
indoc! {"
`log_level` attribute must be one of 'error, warn, info, debug, trace'
= help use `#[log_level(warn)]` or `#[log_level(info)]`
"},
));
}
}
Ok(Some(Ident::new(
level.to_ascii_uppercase().as_str(),
Span::call_site(),
)))
}
fn parse_field_flag(field: &Field) -> FieldFlag {
for attr in field.attrs.iter() {
match &attr.meta {
syn::Meta::Path(path) if path.is_ident("reply") => return FieldFlag::Reply,
_ => {}
}
}
FieldFlag::None
}
fn parse_message_enum(input: DeriveInput) -> Result<Vec<Message>, syn::Error> {
let variants = if let Data::Enum(data_enum) = &input.data {
&data_enum.variants
} else {
return Err(syn::Error::new_spanned(
input,
"handlers can only be derived for enums",
));
};
let mut messages = Vec::new();
for variant in variants {
let name = variant.ident.clone();
let attrs = &variant.attrs;
let message_variant = match &variant.fields {
syn::Fields::Unnamed(fields_) => Variant::Anon {
enum_name: input.ident.clone(),
name,
field_types: fields_
.unnamed
.iter()
.map(|field| field.ty.clone())
.collect(),
field_flags: fields_.unnamed.iter().map(parse_field_flag).collect(),
},
syn::Fields::Named(fields_) => Variant::Named {
enum_name: input.ident.clone(),
name,
field_names: fields_
.named
.iter()
.map(|field| field.ident.clone().unwrap())
.collect(),
field_types: fields_.named.iter().map(|field| field.ty.clone()).collect(),
field_flags: fields_.named.iter().map(parse_field_flag).collect(),
},
_ => {
return Err(syn::Error::new_spanned(
variant,
indoc! {r#"
`Handler` currently only supports named or tuple struct variants
= help use `MyCall(Arg1Type, Arg2Type, ..)`,
= help use `MyCall { arg1: Arg1Type, arg2: Arg2Type, .. }`,
= help use `MyCall(Arg1Type, Arg2Type, .., #[reply] OncePortRef<ReplyType>)`
= help use `MyCall { arg1: Arg1Type, arg2: Arg2Type, .., reply: #[reply] OncePortRef<ReplyType>)`
= help use `MyCall(Arg1Type, Arg2Type, .., #[reply] OncePortHandle<ReplyType>)`
= help use `MyCall { arg1: Arg1Type, arg2: Arg2Type, .., reply: #[reply] OncePortHandle<ReplyType>)`
"#},
));
}
};
let log_level = parse_log_level(attrs)?;
messages.push(Message::new(
variant.fields.span(),
message_variant,
log_level,
)?);
}
Ok(messages)
}
#[proc_macro_derive(Handler, attributes(reply))]
pub fn derive_handler(input: TokenStream) -> TokenStream {
let input = parse_macro_input!(input as DeriveInput);
let name: Ident = input.ident.clone();
let (impl_generics, ty_generics, _) = input.generics.split_for_impl();
let messages = match parse_message_enum(input.clone()) {
Ok(messages) => messages,
Err(err) => return TokenStream::from(err.to_compile_error()),
};
let mut handler_trait_methods = Vec::new();
let mut match_arms = Vec::new();
let mut client_trait_methods = Vec::new();
let global_log_level = parse_log_level(&input.attrs).ok().unwrap_or(None);
for message in messages {
match message {
Message::Call {
ref variant,
ref reply_port_is_handle,
ref return_type,
ref log_level,
} => {
let (arg_names, arg_types): (Vec<_>, Vec<_>) = message.args().into_iter().unzip();
let variant_name_snake = variant.snake_name();
let enum_name = variant.enum_name();
let _variant_qualified_name = variant.qualified_name();
let log_level = match (&global_log_level, log_level) {
(_, Some(local)) => local.clone(),
(Some(global), None) => global.clone(),
_ => Ident::new("DEBUG", Span::call_site()),
};
let _log_level = if *reply_port_is_handle {
quote! {
tracing::Level::#log_level
}
} else {
quote! {
tracing::Level::TRACE
}
};
let log_message = quote! {
hyperactor::metrics::MESSAGES_RECEIVED.add(1, hyperactor::kv_pairs!(
"rpc" => "call",
"actor_id" => this.self_id().to_string(),
"message_type" => stringify!(#enum_name),
"variant" => stringify!(#variant_name_snake),
));
};
handler_trait_methods.push(quote! {
#[doc = "The generated handler method for this enum variant."]
async fn #variant_name_snake(
&mut self,
this: &hyperactor::Instance<Self>,
#(#arg_names: #arg_types),*)
-> Result<#return_type, hyperactor::anyhow::Error>;
});
client_trait_methods.push(quote! {
#[doc = "The generated client method for this enum variant."]
async fn #variant_name_snake(
&self,
caps: &(impl hyperactor::cap::CanSend + hyperactor::cap::CanOpenPort),
#(#arg_names: #arg_types),*)
-> Result<#return_type, hyperactor::anyhow::Error>;
});
let (reply_port_arg, _) = message.reply_port_arg().unwrap();
let constructor = variant.constructor();
let construct_result_future = quote! { use hyperactor::Message; let result = self.#variant_name_snake(this, #(#arg_names),*).await?; };
if *reply_port_is_handle {
match_arms.push(quote! {
#constructor => {
#log_message
#construct_result_future
#reply_port_arg.send(result).map_err(hyperactor::anyhow::Error::from)
}
});
} else {
match_arms.push(quote! {
#constructor => {
#log_message
#construct_result_future
#reply_port_arg.send(this, result).map_err(hyperactor::anyhow::Error::from)
}
});
}
}
Message::OneWay {
ref variant,
ref log_level,
} => {
let (arg_names, arg_types): (Vec<_>, Vec<_>) = message.args().into_iter().unzip();
let variant_name_snake = variant.snake_name();
let enum_name = variant.enum_name();
let log_level = match (&global_log_level, log_level) {
(_, Some(local)) => local.clone(),
(Some(global), None) => global.clone(),
_ => Ident::new("TRACE", Span::call_site()),
};
let _log_level = quote! {
tracing::Level::#log_level
};
let log_message = quote! {
hyperactor::metrics::MESSAGES_RECEIVED.add(1, hyperactor::kv_pairs!(
"rpc" => "call",
"actor_id" => this.self_id().to_string(),
"message_type" => stringify!(#enum_name),
"variant" => stringify!(#variant_name_snake),
));
};
handler_trait_methods.push(quote! {
#[doc = "The generated handler method for this enum variant."]
async fn #variant_name_snake(
&mut self,
this: &hyperactor::Instance<Self>,
#(#arg_names: #arg_types),*)
-> Result<(), hyperactor::anyhow::Error>;
});
client_trait_methods.push(quote! {
#[doc = "The generated client method for this enum variant."]
async fn #variant_name_snake(
&self,
caps: &impl hyperactor::cap::CanSend,
#(#arg_names: #arg_types),*)
-> Result<(), hyperactor::anyhow::Error>;
});
let constructor = variant.constructor();
match_arms.push(quote! {
#constructor => {
#log_message
self.#variant_name_snake(this, #(#arg_names),*).await
},
});
}
}
}
let handler_trait_name = format_ident!("{}Handler", name);
let client_trait_name = format_ident!("{}Client", name);
let expanded = quote! {
#[doc = "The custom handler trait for this message type."]
#[hyperactor::async_trait::async_trait]
pub trait #handler_trait_name #impl_generics: hyperactor::Actor + Send + Sync {
#(#handler_trait_methods)*
#[doc = "Handle the next message."]
async fn handle(
&mut self,
this: &hyperactor::Instance<Self>,
message: #name #ty_generics,
) -> hyperactor::anyhow::Result<()> {
match message {
#(#match_arms)*
}
}
}
#[doc = "The custom client trait for this message type."]
#[hyperactor::async_trait::async_trait]
pub trait #client_trait_name #impl_generics: Send + Sync {
#(#client_trait_methods)*
}
};
TokenStream::from(expanded)
}
#[proc_macro_derive(HandleClient, attributes(log_level))]
pub fn derive_handle_client(input: TokenStream) -> TokenStream {
derive_client(input, true)
}
#[proc_macro_derive(RefClient, attributes(log_level))]
pub fn derive_ref_client(input: TokenStream) -> TokenStream {
derive_client(input, false)
}
fn derive_client(input: TokenStream, is_handle: bool) -> TokenStream {
let input = parse_macro_input!(input as DeriveInput);
let name = input.ident.clone();
let messages = match parse_message_enum(input.clone()) {
Ok(messages) => messages,
Err(err) => return TokenStream::from(err.to_compile_error()),
};
let mut impl_methods = Vec::new();
let send_message = if is_handle {
quote! { self.send(message)? }
} else {
quote! { self.send(caps, message)? }
};
let global_log_level = parse_log_level(&input.attrs).ok().unwrap_or(None);
for message in messages {
match message {
Message::Call {
ref variant,
ref reply_port_is_handle,
ref return_type,
ref log_level,
} => {
let (arg_names, arg_types): (Vec<_>, Vec<_>) = message.args().into_iter().unzip();
let variant_name_snake = variant.snake_name();
let enum_name = variant.enum_name();
let (reply_port_arg, _) = message.reply_port_arg().unwrap();
let constructor = variant.constructor();
let log_level = match (&global_log_level, log_level) {
(_, Some(local)) => local.clone(),
(Some(global), None) => global.clone(),
_ => Ident::new("DEBUG", Span::call_site()),
};
let log_level = if is_handle {
quote! {
tracing::Level::#log_level
}
} else {
quote! {
tracing::Level::TRACE
}
};
let log_message = quote! {
hyperactor::metrics::MESSAGES_SENT.add(1, hyperactor::kv_pairs!(
"rpc" => "call",
"actor_id" => self.actor_id().to_string(),
"message_type" => stringify!(#enum_name),
"variant" => stringify!(#variant_name_snake),
));
tracing::event!(target: "message", #log_level, rpc = "call", payload=?message, "send");
};
if *reply_port_is_handle {
impl_methods.push(quote! {
#[hyperactor::instrument(level=#log_level, rpc = "call", message_type=#name)]
async fn #variant_name_snake(
&self,
caps: &(impl hyperactor::cap::CanSend + hyperactor::cap::CanOpenPort),
#(#arg_names: #arg_types),*)
-> Result<#return_type, hyperactor::anyhow::Error> {
let (#reply_port_arg, reply_receiver) =
hyperactor::mailbox::open_once_port::<#return_type>(caps);
let message = #constructor;
#log_message;
#send_message;
reply_receiver.recv().await.map_err(hyperactor::anyhow::Error::from)
}
});
} else {
impl_methods.push(quote! {
#[hyperactor::instrument(level=#log_level, rpc="call", message_type=#name)]
async fn #variant_name_snake(
&self,
caps: &(impl hyperactor::cap::CanSend + hyperactor::cap::CanOpenPort),
#(#arg_names: #arg_types),*)
-> Result<#return_type, hyperactor::anyhow::Error> {
let (#reply_port_arg, reply_receiver) =
hyperactor::mailbox::open_once_port::<#return_type>(caps);
let #reply_port_arg = #reply_port_arg.bind();
let message = #constructor;
#log_message;
#send_message;
reply_receiver.recv().await.map_err(hyperactor::anyhow::Error::from)
}
});
}
}
Message::OneWay {
ref variant,
ref log_level,
} => {
let (arg_names, arg_types): (Vec<_>, Vec<_>) = message.args().into_iter().unzip();
let variant_name_snake = variant.snake_name();
let enum_name = variant.enum_name();
let constructor = variant.constructor();
let log_level = match (&global_log_level, log_level) {
(_, Some(local)) => local.clone(),
(Some(global), None) => global.clone(),
_ => Ident::new("DEBUG", Span::call_site()),
};
let log_level = if is_handle {
quote! {
tracing::Level::TRACE
}
} else {
quote! {
tracing::Level::#log_level
}
};
let log_message = quote! {
hyperactor::metrics::MESSAGES_SENT.add(1, hyperactor::kv_pairs!(
"rpc" => "oneway",
"actor_id" => self.actor_id().to_string(),
"message_type" => stringify!(#enum_name),
"variant" => stringify!(#variant_name_snake),
));
tracing::event!(target: "message", #log_level, handle = stringify!(#variant_name_snake), rpc = "oneway", payload=?message, "send");
};
impl_methods.push(quote! {
async fn #variant_name_snake(
&self,
caps: &impl hyperactor::cap::CanSend,
#(#arg_names: #arg_types),*)
-> Result<(), hyperactor::anyhow::Error> {
let message = #constructor;
#log_message;
#send_message;
Ok(())
}
});
}
}
}
let trait_name = format_ident!("{}Client", name);
let (_, ty_generics, _) = input.generics.split_for_impl();
let a_ident = Ident::new("A", proc_macro2::Span::from(proc_macro::Span::def_site()));
let mut trait_generics = input.generics.clone();
trait_generics.params.insert(
0,
syn::GenericParam::Type(syn::TypeParam {
ident: a_ident.clone(),
attrs: vec![],
colon_token: None,
bounds: Punctuated::new(),
eq_token: None,
default: None,
}),
);
let (impl_generics, _, _) = trait_generics.split_for_impl();
let expanded = if is_handle {
quote! {
#[hyperactor::async_trait::async_trait]
impl #impl_generics #trait_name #ty_generics for hyperactor::ActorHandle<#a_ident>
where #a_ident: hyperactor::Handler<#name #ty_generics> {
#(#impl_methods)*
}
}
} else {
quote! {
#[hyperactor::async_trait::async_trait]
impl #impl_generics #trait_name #ty_generics for hyperactor::ActorRef<#a_ident>
where #a_ident: hyperactor::actor::RemoteHandles<#name #ty_generics> {
#(#impl_methods)*
}
}
};
TokenStream::from(expanded)
}
const FORWARD_ARGUMENT_ERROR: &str = indoc! {r#"
`forward` expects the message type that is being forwarded
= help: use `#[forward(MessageType)]`
"#};
#[proc_macro_attribute]
pub fn forward(attr: TokenStream, item: TokenStream) -> TokenStream {
let attr_args = parse_macro_input!(attr with Punctuated::<syn::PathSegment, syn::Token![,]>::parse_terminated);
if attr_args.len() != 1 {
return TokenStream::from(
syn::Error::new_spanned(attr_args, FORWARD_ARGUMENT_ERROR).to_compile_error(),
);
}
let message_type = attr_args.first().unwrap();
let input = parse_macro_input!(item as ItemImpl);
let self_type = match *input.self_ty {
syn::Type::Path(ref type_path) => {
let segment = type_path.path.segments.last().unwrap();
segment.clone() }
_ => {
return TokenStream::from(
syn::Error::new_spanned(input.self_ty, "`forward` argument must be a type")
.to_compile_error(),
);
}
};
let trait_name = match input.trait_ {
Some((_, ref trait_path, _)) => trait_path.segments.last().unwrap().clone(),
None => {
return TokenStream::from(
syn::Error::new_spanned(input.self_ty, "no trait in implementation block")
.to_compile_error(),
);
}
};
let expanded = quote! {
#input
#[hyperactor::async_trait::async_trait]
impl hyperactor::Handler<#message_type> for #self_type {
async fn handle(
&mut self,
this: &hyperactor::Instance<Self>,
message: #message_type,
) -> hyperactor::anyhow::Result<()> {
<Self as #trait_name>::handle(self, this, message).await
}
}
};
TokenStream::from(expanded)
}
#[proc_macro_attribute]
pub fn instrument(args: TokenStream, input: TokenStream) -> TokenStream {
let args =
parse_macro_input!(args with Punctuated::<syn::Expr, syn::Token![,]>::parse_terminated);
let input = parse_macro_input!(input as ItemFn);
let output = quote! {
#[hyperactor::tracing::instrument(err, skip_all, #args)]
#input
};
TokenStream::from(output)
}
#[proc_macro_attribute]
pub fn instrument_infallible(args: TokenStream, input: TokenStream) -> TokenStream {
let args =
parse_macro_input!(args with Punctuated::<syn::Expr, syn::Token![,]>::parse_terminated);
let input = parse_macro_input!(input as ItemFn);
let output = quote! {
#[hyperactor::tracing::instrument(skip_all, #args)]
#input
};
TokenStream::from(output)
}
#[proc_macro_derive(Named, attributes(named))]
pub fn named_derive(input: TokenStream) -> TokenStream {
let input = parse_macro_input!(input as DeriveInput);
let struct_name = &input.ident;
let mut typename = quote! {
concat!(std::module_path!(), "::", stringify!(#struct_name))
};
let mut dump = true;
for attr in &input.attrs {
if attr.path().is_ident("named") {
if let Ok(meta) = attr.parse_args_with(
syn::punctuated::Punctuated::<Meta, syn::Token![,]>::parse_terminated,
) {
for item in meta {
if let Meta::NameValue(MetaNameValue {
path,
value: Expr::Lit(expr_lit),
..
}) = item
{
if path.is_ident("dump") {
if let Lit::Bool(lit_bool) = expr_lit.lit {
dump = lit_bool.value;
}
} else if path.is_ident("name") {
if let Lit::Str(name) = expr_lit.lit {
typename = quote! { #name };
}
}
}
}
}
}
}
let cached_typehash = Ident::new(
&format!("{}_CACHED_TYPEHASH", struct_name).to_case(Case::UpperSnake),
Span::call_site(),
);
let dumper = if dump {
quote! { Some(<#struct_name as hyperactor::data::NamedDumpable>::dump) }
} else {
quote! { None }
};
let arm_impl = match &input.data {
Data::Enum(DataEnum { variants, .. }) => {
let match_arms = variants.iter().map(|v| {
let variant_name = &v.ident;
let variant_str = variant_name.to_string();
match &v.fields {
Fields::Unit => quote! { Self::#variant_name => Some(#variant_str) },
Fields::Unnamed(_) => quote! { Self::#variant_name(..) => Some(#variant_str) },
Fields::Named(_) => quote! { Self::#variant_name { .. } => Some(#variant_str) },
}
});
quote! {
fn arm(&self) -> Option<&'static str> {
match self {
#(#match_arms,)*
}
}
}
}
_ => quote! {},
};
let expanded = quote! {
static #cached_typehash: std::sync::LazyLock<u64> = std::sync::LazyLock::new(|| {
hyperactor::cityhasher::hash(<#struct_name as hyperactor::data::Named>::typename())
});
impl hyperactor::data::Named for #struct_name {
fn typename() -> &'static str { #typename }
fn typehash() -> u64 { *#cached_typehash }
#arm_impl
}
hyperactor::submit! {
hyperactor::data::TypeInfo {
typename: <#struct_name as hyperactor::data::Named>::typename,
typehash: <#struct_name as hyperactor::data::Named>::typehash,
typeid: <#struct_name as hyperactor::data::Named>::typeid,
port: <#struct_name as hyperactor::data::Named>::port,
dump: #dumper,
arm_unchecked: <#struct_name as hyperactor::data::Named>::arm_unchecked,
}
}
};
TokenStream::from(expanded)
}
#[proc_macro_attribute]
pub fn export(attr: TokenStream, item: TokenStream) -> TokenStream {
export_impl("export", attr, &parse_macro_input!(item as DeriveInput))
}
#[proc_macro_attribute]
pub fn export_spawn(attr: TokenStream, item: TokenStream) -> TokenStream {
let input: DeriveInput = parse_macro_input!(item as DeriveInput);
let mut exported = export_impl("export_spawn", attr, &input);
let data_type_name = &input.ident;
exported.extend(TokenStream::from(quote! {
hyperactor::remote!(#data_type_name);
}));
exported
}
fn export_impl(which: &'static str, attr: TokenStream, input: &DeriveInput) -> TokenStream {
let data_type_name = &input.ident;
let attr_args =
parse_macro_input!(attr with Punctuated::<syn::Type, Token![,]>::parse_terminated);
if attr_args.is_empty() {
return TokenStream::from(
syn::Error::new_spanned(attr_args, format!("`{}` expects one or more type path arguments\n\n= help: use `#[{}(MyType, MyOtherType<T>)]`", which, which)).to_compile_error(),
);
}
let mut handles = Vec::new();
let mut bindings = Vec::new();
for ty in &attr_args {
handles.push(quote! {
impl hyperactor::actor::RemoteHandles<#ty> for #data_type_name {}
});
bindings.push(quote! {
ports.bind::<#ty>();
});
}
let expanded = quote! {
#input
impl hyperactor::actor::RemoteActor for #data_type_name {}
#(#handles)*
impl hyperactor::actor::RemoteHandles<hyperactor::actor::Signal> for #data_type_name {}
impl hyperactor::actor::Binds<#data_type_name> for #data_type_name {
fn bind(ports: &hyperactor::proc::Ports<Self>) {
#(#bindings)*
}
}
impl hyperactor::data::Named for #data_type_name {
fn typename() -> &'static str { concat!(std::module_path!(), "::", stringify!(#data_type_name)) }
}
};
TokenStream::from(expanded)
}