use proc_macro::TokenStream;
use quote::quote;
use syn::{parse_macro_input, Attribute, FnArg, ItemTrait, ReturnType, TraitItem, Type};
fn parse_rpc_args(args: &str) -> RpcConfig {
let server = args.contains("server");
let client = args.contains("client");
let mut config = RpcConfig {
server,
client,
..Default::default()
};
if let Some(start) = args.find("namespace = \"") {
let start = start + 13;
if let Some(end) = args[start..].find('"') {
config.namespace = args[start..start + end].to_string();
}
}
config
}
fn extract_result_inner_type(return_type: Option<&Type>) -> proc_macro2::TokenStream {
match return_type {
Some(Type::Path(type_path)) => {
if let Some(segment) = type_path.path.segments.last() {
if segment.ident == "Result" {
match &segment.arguments {
syn::PathArguments::AngleBracketed(args) => {
if let Some(syn::GenericArgument::Type(inner_type)) = args.args.first()
{
quote! { #inner_type }
} else {
quote! { () }
}
}
_ => quote! { () },
}
} else {
quote! { #return_type }
}
} else {
quote! { () }
}
}
_ => quote! { () },
}
}
#[derive(Default)]
#[allow(dead_code)]
struct RpcConfig {
server: bool,
client: bool,
namespace: String,
}
#[derive(Debug, PartialEq)]
enum MethodType {
Method, Subscription, }
fn parse_method_attributes(attrs: &[Attribute], default_name: &str) -> (MethodType, String) {
for attr in attrs {
if attr.path().is_ident("method") {
if let Ok(meta) = attr.meta.require_list() {
let tokens = meta.tokens.to_string();
if let Some(start) = tokens.find("name = \"") {
let start = start + 8; if let Some(end) = tokens[start..].find('"') {
let method_name = tokens[start..start + end].to_string();
return (MethodType::Method, method_name);
}
}
}
return (MethodType::Method, default_name.to_string());
} else if attr.path().is_ident("subscription") {
if let Ok(meta) = attr.meta.require_list() {
let tokens = meta.tokens.to_string();
if let Some(start) = tokens.find("name = \"") {
let start = start + 8; if let Some(end) = tokens[start..].find('"') {
let subscription_name = tokens[start..start + end].to_string();
return (MethodType::Subscription, subscription_name);
}
}
}
return (MethodType::Subscription, default_name.to_string());
}
}
(MethodType::Method, default_name.to_string())
}
fn generate_subscription_handler(
_method_name: &syn::Ident,
rpc_method_name: &str,
) -> proc_macro2::TokenStream {
quote! {
#rpc_method_name => {
Err(hsipc::Error::method_not_found(self.name(), method))
}
}
}
fn generate_method_handler(
method_name: &syn::Ident,
rpc_method_name: &str,
params: &[&Type],
is_async: bool,
) -> proc_macro2::TokenStream {
if params.len() == 1 {
let param_type = params[0];
if is_async {
quote! {
#rpc_method_name => {
let request: #param_type = bincode::deserialize(&payload)?;
let response = self.inner.#method_name(request).await?;
Ok(bincode::serialize(&response)?)
}
}
} else {
quote! {
#rpc_method_name => {
let request: #param_type = bincode::deserialize(&payload)?;
let response = self.inner.#method_name(request)?;
Ok(bincode::serialize(&response)?)
}
}
}
} else if params.is_empty() {
if is_async {
quote! {
#rpc_method_name => {
let response = self.inner.#method_name().await?;
Ok(bincode::serialize(&response)?)
}
}
} else {
quote! {
#rpc_method_name => {
let response = self.inner.#method_name()?;
Ok(bincode::serialize(&response)?)
}
}
}
} else {
let param_tuple = quote! { (#(#params),*) };
if is_async {
quote! {
#rpc_method_name => {
let params: #param_tuple = bincode::deserialize(&payload)?;
let response = self.inner.#method_name(params.0, params.1).await?;
Ok(bincode::serialize(&response)?)
}
}
} else {
quote! {
#rpc_method_name => {
let params: #param_tuple = bincode::deserialize(&payload)?;
let response = self.inner.#method_name(params.0, params.1)?;
Ok(bincode::serialize(&response)?)
}
}
}
}
}
fn generate_rpc_client_method(
method_name: &syn::Ident,
rpc_method_name: &str,
params: &[&Type],
client_return_type: &proc_macro2::TokenStream,
namespace: &str,
is_async: bool,
) -> proc_macro2::TokenStream {
if params.len() == 1 {
let param_type = params[0];
if is_async {
quote! {
pub async fn #method_name(&self, request: #param_type) -> hsipc::Result<#client_return_type> {
let result: #client_return_type = self.hub.call(&format!("{}.{}", #namespace, #rpc_method_name), request).await?;
Ok(result)
}
}
} else {
quote! {
pub fn #method_name(&self, request: #param_type) -> hsipc::Result<#client_return_type> {
let result: #client_return_type = futures::executor::block_on(
self.hub.call(&format!("{}.{}", #namespace, #rpc_method_name), request)
)?;
Ok(result)
}
}
}
} else if params.is_empty() {
if is_async {
quote! {
pub async fn #method_name(&self) -> hsipc::Result<#client_return_type> {
let result: #client_return_type = self.hub.call(&format!("{}.{}", #namespace, #rpc_method_name), ()).await?;
Ok(result)
}
}
} else {
quote! {
pub fn #method_name(&self) -> hsipc::Result<#client_return_type> {
let result: #client_return_type = futures::executor::block_on(
self.hub.call(&format!("{}.{}", #namespace, #rpc_method_name), ())
)?;
Ok(result)
}
}
}
} else {
let param_names: Vec<syn::Ident> = (0..params.len())
.map(|i| syn::Ident::new(&format!("p{i}"), method_name.span()))
.collect();
if is_async {
quote! {
pub async fn #method_name(&self, #(#param_names: #params),*) -> hsipc::Result<#client_return_type> {
let params = (#(#param_names),*);
let result: #client_return_type = self.hub.call(&format!("{}.{}", #namespace, #rpc_method_name), params).await?;
Ok(result)
}
}
} else {
quote! {
pub fn #method_name(&self, #(#param_names: #params),*) -> hsipc::Result<#client_return_type> {
let params = (#(#param_names),*);
let result: #client_return_type = futures::executor::block_on(
self.hub.call(&format!("{}.{}", #namespace, #rpc_method_name), params)
)?;
Ok(result)
}
}
}
}
}
fn transform_trait_for_subscription(input: &ItemTrait) -> proc_macro2::TokenStream {
let trait_ident = &input.ident;
let trait_generics = &input.generics;
let trait_bounds = &input.supertraits;
let mut transformed_items = Vec::new();
for item in &input.items {
if let TraitItem::Fn(method) = item {
let method_name = &method.sig.ident;
let method_name_str = method_name.to_string();
let (method_type, _) = parse_method_attributes(&method.attrs, &method_name_str);
if method_type == MethodType::Subscription {
let mut transformed_method = method.clone();
let mut new_inputs = syn::punctuated::Punctuated::new();
if let Some(first_input) = transformed_method.sig.inputs.first() {
new_inputs.push(first_input.clone());
}
let pending_param: syn::FnArg =
syn::parse_str("pending: hsipc::PendingSubscriptionSink").unwrap();
new_inputs.push(pending_param);
for input in transformed_method.sig.inputs.iter().skip(1) {
new_inputs.push(input.clone());
}
transformed_method.sig.inputs = new_inputs;
transformed_items.push(TraitItem::Fn(transformed_method));
} else {
transformed_items.push(item.clone());
}
} else {
transformed_items.push(item.clone());
}
}
quote! {
pub trait #trait_ident #trait_generics: #trait_bounds {
#(#transformed_items)*
}
}
}
fn generate_subscription_client_method(
method_name: &syn::Ident,
rpc_method_name: &str,
params: &[&Type],
namespace: &str,
_return_type: Option<&Type>,
) -> proc_macro2::TokenStream {
if params.len() == 1 {
let param_type = params[0];
quote! {
pub async fn #method_name(&self, params: #param_type) -> hsipc::Result<()> {
let serialized_params = bincode::serialize(¶ms)?;
let request_msg = hsipc::Message::subscription_request(
self.hub.name().to_string(),
None, format!("{}.{}", #namespace, #rpc_method_name),
serialized_params,
);
Ok(())
}
}
} else if params.is_empty() {
quote! {
pub async fn #method_name(&self) -> hsipc::Result<()> {
let request_msg = hsipc::Message::subscription_request(
self.hub.name().to_string(),
None, format!("{}.{}", #namespace, #rpc_method_name),
vec![], );
Ok(())
}
}
} else {
let param_names: Vec<syn::Ident> = (0..params.len())
.map(|i| syn::Ident::new(&format!("p{i}"), method_name.span()))
.collect();
quote! {
pub async fn #method_name(&self, #(#param_names: #params),*) -> hsipc::Result<()> {
let params_tuple = (#(#param_names),*);
let serialized_params = bincode::serialize(¶ms_tuple)?;
let request_msg = hsipc::Message::subscription_request(
self.hub.name().to_string(),
None, format!("{}.{}", #namespace, #rpc_method_name),
serialized_params,
);
Ok(())
}
}
}
}
pub fn rpc_impl(args: TokenStream, input: TokenStream) -> TokenStream {
let input = parse_macro_input!(input as ItemTrait);
let args_str = args.to_string();
let config = parse_rpc_args(&args_str);
let trait_name = &input.ident;
let service_name = syn::Ident::new(&format!("{trait_name}Service"), trait_name.span());
let client_name = syn::Ident::new(&format!("{trait_name}Client"), trait_name.span());
let namespace = &config.namespace;
let mut method_names = Vec::new();
let mut service_handlers = Vec::new();
let mut client_methods = Vec::new();
for item in &input.items {
if let TraitItem::Fn(method) = item {
let method_name = &method.sig.ident;
let method_name_str = method_name.to_string();
let (method_type, rpc_method_name) =
parse_method_attributes(&method.attrs, &method_name_str);
method_names.push(rpc_method_name.clone());
let params: Vec<&Type> = method
.sig
.inputs
.iter()
.filter_map(|arg| match arg {
FnArg::Typed(pat_type) => Some(&*pat_type.ty),
_ => None,
})
.collect();
let return_type = match &method.sig.output {
ReturnType::Type(_, ty) => Some(&**ty),
ReturnType::Default => None,
};
let is_async = method.sig.asyncness.is_some();
let handler = match method_type {
MethodType::Subscription => {
generate_subscription_handler(method_name, &rpc_method_name)
}
MethodType::Method => {
generate_method_handler(method_name, &rpc_method_name, ¶ms, is_async)
}
};
service_handlers.push(handler);
let client_method = match method_type {
MethodType::Subscription => {
generate_subscription_client_method(
method_name,
&rpc_method_name,
¶ms,
namespace,
return_type,
)
}
MethodType::Method => {
let client_return_type = extract_result_inner_type(return_type);
generate_rpc_client_method(
method_name,
&rpc_method_name,
¶ms,
&client_return_type,
namespace,
is_async,
)
}
};
client_methods.push(client_method);
}
}
let transformed_trait = transform_trait_for_subscription(&input);
let expanded = quote! {
#[hsipc::async_trait]
#transformed_trait
pub struct #service_name<T> {
inner: T,
}
impl<T> #service_name<T>
where
T: #trait_name + Send + Sync,
{
pub fn new(inner: T) -> Self {
Self { inner }
}
}
#[hsipc::async_trait]
impl<T> hsipc::Service for #service_name<T>
where
T: #trait_name + Send + Sync + 'static,
{
fn name(&self) -> &'static str {
#namespace
}
fn methods(&self) -> Vec<&'static str> {
vec![#(#method_names),*]
}
async fn handle(&self, method: &str, payload: Vec<u8>) -> hsipc::Result<Vec<u8>> {
match method {
#(#service_handlers)*
_ => Err(hsipc::Error::method_not_found(self.name(), method))
}
}
}
#[derive(Clone)]
pub struct #client_name {
hub: hsipc::ProcessHub,
}
impl #client_name {
pub fn new(hub: hsipc::ProcessHub) -> Self {
Self { hub }
}
#(#client_methods)*
}
};
TokenStream::from(expanded)
}