jsonrpc-server-macro 0.1.0

jsonrpc-server-macro
Documentation
use proc_macro::TokenStream;
use quote::{quote, quote_spanned};
use syn::punctuated::Punctuated;
use syn::spanned::Spanned;
use syn::{
    Error, FnArg, ImplItem, ImplItemFn, Item, ItemFn, ItemImpl, Pat, ReturnType, Signature, Token,
    Type, parse, parse_macro_input,
};

#[proc_macro_attribute]
pub fn jsonrpc(_attr: TokenStream, input: TokenStream) -> TokenStream {
    match parse_macro_input!(input as Item) {
        Item::Fn(item) => expand_fn(item).unwrap_or_else(|e| e.to_compile_error().into()),
        Item::Impl(item) => expand_impl(item).unwrap_or_else(|e| e.to_compile_error().into()),
        item => Error::new_spanned(item, "#[jsonrpc]: expected fn or impl block")
            .to_compile_error()
            .into(),
    }
}

fn expand_impl(item: ItemImpl) -> Result<TokenStream, Error> {
    if let Some((_, ref path, ..)) = item.trait_ {
        return Err(Error::new_spanned(
            path,
            "#[jsonrpc]: trait impl is not supported",
        ));
    }

    if !item.generics.params.is_empty() {
        return Err(Error::new_spanned(
            item.generics,
            "#[jsonrpc]: generic is not supported",
        ));
    }

    let prefix = match *item.self_ty {
        Type::Path(ref path) => snake_case(
            path.path
                .segments
                .last()
                .unwrap()
                .ident
                .to_string()
                .as_bytes(),
        ),
        _ => {
            return Err(Error::new_spanned(
                item.self_ty,
                "#[jsonrpc]: not supported",
            ));
        }
    };

    let mut names = Vec::new();
    let mut methods = Vec::new();
    for impl_item in &item.items {
        if let ImplItem::Fn(item_fn) = impl_item {
            names.push(format!("{}.{}", prefix, item_fn.sig.ident));
            methods.push(generate_method(&item, item_fn)?);
        }
    }

    let ty = &item.self_ty;
    Ok(quote! {
        #item
        impl jsonrpc_server::Register for #ty {
            fn register(&self, registry: &mut jsonrpc_server::Registry) {
                #(registry.add(#names, #methods);)*
            }
        }
    }
    .into())
}

fn generate_method(
    item: &ItemImpl,
    item_fn: &ImplItemFn,
) -> Result<proc_macro2::TokenStream, Error> {
    if !item_fn.sig.generics.params.is_empty() {
        return Err(Error::new_spanned(
            &item_fn.sig.generics,
            "#[jsonrpc]: generic is not supported",
        ));
    }
    let ident = &item_fn.sig.ident;
    let self_ty = &item.self_ty;
    let ret_assert = ret_assert(&item_fn.sig)?;
    let (arg_assert, args) = arg_assert(&item_fn.sig.inputs, true)?;
    let wait = item_fn
        .sig
        .asyncness
        .map(|_| quote!(let result = result.await;));
    let set_method = set_method(ident);
    let argc = 0..args.len();
    let st = quote_spanned! {item_fn.span()=>
        {
            struct __Method(#self_ty);

            impl jsonrpc_server::Method for __Method {
                fn call(&self, args: jsonrpc_server::serde_json::Value) -> jsonrpc_server::BoxFuture<'_> {
                    #arg_assert
                    #ret_assert

                    #[allow(unused)]
                    macro_rules! arg {
                        ($v:expr) => {
                            jsonrpc_server::serde_json::from_value($v).map_err(|err| {
                                jsonrpc_server::error!("deserialize parameter error: {}", err);
                                jsonrpc_server::Error::invalid_params()
                            })
                        };
                    }

                    Box::pin(async move {
                        #[allow(unused)]
                        let result = match args {
                            jsonrpc_server::serde_json::Value::Array(mut args) => {
                                self.0.#ident(#(arg!(args.get_mut(#argc).map(jsonrpc_server::serde_json::Value::take).unwrap_or(jsonrpc_server::serde_json::Value::Null))?),*)
                            }
                            jsonrpc_server::serde_json::Value::Object(mut args) => {
                                self.0.#ident(#(arg!(args.remove(#args).unwrap_or(jsonrpc_server::serde_json::Value::Null))?),*)
                            }
                            _ => return Err(jsonrpc_server::Error::invalid_params()),
                        };
                        #wait
                        #set_method
                        Ok(jsonrpc_server::serde_json::to_value(result?).expect("serialize error"))
                    })
                }
            }

            __Method(self.clone())
        }
    };
    parse(st.into())
}

fn expand_fn(item: ItemFn) -> Result<TokenStream, Error> {
    let ident = &item.sig.ident;
    let vis = &item.vis;
    if !item.sig.generics.params.is_empty() {
        return Err(Error::new_spanned(
            item.sig.generics,
            "#[jsonrpc]: generic is not supported",
        ));
    }

    let ret_assert = ret_assert(&item.sig)?;
    let (arg_assert, args) = arg_assert(&item.sig.inputs, false)?;
    let wait = item
        .sig
        .asyncness
        .map(|_| quote!(let result = result.await;));
    let set_method = set_method(ident);
    let argc = 0..args.len();
    let ts = quote! {
        #vis fn #ident(args: jsonrpc_server::serde_json::Value) -> std::pin::Pin<Box<dyn std::future::Future<Output=std::result::Result<jsonrpc_server::serde_json::Value, jsonrpc_server::Error>> + Send>> {
            #arg_assert
            #ret_assert
            #item

            #[allow(unused)]
            macro_rules! arg {
                ($v:expr) => {
                    jsonrpc_server::serde_json::from_value($v).map_err(|err| {
                        jsonrpc_server::error!("deserialize parameter error: {}", err);
                        jsonrpc_server::Error::invalid_params()
                    })
                };
            }

            Box::pin(async move {
                #[allow(unused)]
                let result = match args {
                    jsonrpc_server::serde_json::Value::Array(mut args) => {
                        #ident(#(arg!(args.get_mut(#argc).map(jsonrpc_server::serde_json::Value::take).unwrap_or(jsonrpc_server::serde_json::Value::Null))?),*)
                    }
                    jsonrpc_server::serde_json::Value::Object(mut args) => {
                        #ident(#(arg!(args.remove(#args).unwrap_or(jsonrpc_server::serde_json::Value::Null))?),*)
                    }
                    _ => return Err(jsonrpc_server::Error::invalid_params()),
                };
                #wait
                #set_method
                Ok(jsonrpc_server::serde_json::to_value(result?).expect("serialize error"))
            })
        }
    };
    Ok(ts.into())
}

#[cfg(feature = "anyhow")]
fn set_method(name: &proc_macro2::Ident) -> Option<proc_macro2::TokenStream> {
    let name = name.to_string();
    Some(quote! {
        let result = match result {
            Ok(v) => Ok(v),
            Err(e) => {
                let mut e = jsonrpc_server::Error::from(e);
                e.method = Some(#name);
                Err(e)
            }
        };
    })
}

#[cfg(not(feature = "anyhow"))]
fn set_method(_: &proc_macro2::Ident) -> Option<proc_macro2::TokenStream> {
    None
}

fn ret_assert(sig: &Signature) -> Result<proc_macro2::TokenStream, Error> {
    match sig.output {
        ReturnType::Default => Err(Error::new_spanned(sig, "#[jsonrpc]: expected return value")),
        ReturnType::Type(_, ref ty) => Ok(quote_spanned! {ty.span()=>
            {
                fn assert(_: Option<std::result::Result<impl jsonrpc_server::serde::Serialize, impl Into<jsonrpc_server::Error>>>) {}
                assert(None::<#ty>);
            }
        }),
    }
}

fn arg_assert(
    inputs: &Punctuated<FnArg, Token![,]>,
    is_method: bool,
) -> Result<(proc_macro2::TokenStream, Vec<String>), Error> {
    let mut assert = vec![];
    let mut args = Vec::with_capacity(inputs.len());
    for arg in inputs.iter().skip(if is_method { 1 } else { 0 }) {
        match arg {
            FnArg::Typed(arg) => match *arg.pat {
                Pat::Ident(ref pat) => {
                    args.push(pat.ident.to_string());
                    let ty = &arg.ty;
                    assert.push(quote_spanned! {ty.span()=>
                        { struct _Assert where #ty: jsonrpc_server::serde::de::DeserializeOwned; }
                    })
                }
                _ => return Err(Error::new_spanned(arg, "#[jsonrpc]: unsupported argument")),
            },
            FnArg::Receiver(_) => unreachable!(),
        }
    }
    Ok((quote!(#(#assert)*), args))
}

fn snake_case(s: &[u8]) -> String {
    let mut result = String::with_capacity(s.len());
    for &b in s {
        match b {
            b'A'..=b'Z' => {
                if !result.is_empty() {
                    result.push('_');
                }
                result.push((b + 32) as char);
            }
            b => result.push(b as char),
        }
    }
    result
}