coreum_std_derive/
lib.rs

1use itertools::Itertools;
2use proc_macro::TokenStream;
3use proc_macro2::TokenTree;
4use quote::quote;
5use syn::{parse_macro_input, DeriveInput};
6
7macro_rules! match_kv_attr {
8    ($key:expr, $value_type:tt) => {
9        |tt| {
10            if let [TokenTree::Ident(key), TokenTree::Punct(eq), TokenTree::$value_type(value)] =
11                &tt[..]
12            {
13                if (key == $key) && (eq.as_char() == '=') {
14                    Some(quote!(#value))
15                } else {
16                    None
17                }
18            } else {
19                None
20            }
21        }
22    };
23}
24
25#[proc_macro_derive(CosmwasmExt, attributes(proto_message, proto_query))]
26pub fn derive_cosmwasm_ext(input: TokenStream) -> TokenStream {
27    let input = parse_macro_input!(input as DeriveInput);
28    let ident = input.ident;
29
30    let type_url = get_type_url(&input.attrs);
31
32    // `EncodeError` always indicates that a message failed to encode because the
33    // provided buffer had insufficient capacity. Message encoding is otherwise
34    // infallible.
35
36    let (query_request_conversion, cosmwasm_query) = if get_attr("proto_query", &input.attrs)
37        .is_some()
38    {
39        let path = get_query_attrs(&input.attrs, match_kv_attr!("path", Literal));
40        let res = get_query_attrs(&input.attrs, match_kv_attr!("response_type", Ident));
41
42        let query_request_conversion = quote! {
43            impl From<#ident> for cosmwasm_std::CosmosMsg {
44                fn from(msg: #ident) -> Self {
45                    cosmwasm_std::CosmosMsg::Any(cosmwasm_std::AnyMsg {
46                        type_url: #path.to_string(),
47                        value: cosmwasm_std::Binary::from(msg.to_proto_bytes()),
48                    })
49                }
50            }
51        };
52
53        let cosmwasm_query = quote! {
54            pub fn query(self, querier: &cosmwasm_std::QuerierWrapper<impl cosmwasm_std::CustomQuery>) -> cosmwasm_std::StdResult<#res> {
55                Ok(#res::try_from(querier.query_grpc(#path.to_string(), self.into())?).unwrap())
56            }
57        };
58
59        (query_request_conversion, cosmwasm_query)
60    } else {
61        (quote!(), quote!())
62    };
63
64    (quote! {
65        impl #ident {
66            pub const TYPE_URL: &'static str = #type_url;
67            #cosmwasm_query
68
69            pub fn to_proto_bytes(&self) -> Vec<u8> {
70                let mut bytes = Vec::new();
71                prost::Message::encode(self, &mut bytes)
72                    .expect("Message encoding must be infallible");
73                bytes
74            }
75            pub fn to_any(&self) -> cosmwasm_std::AnyMsg {
76                cosmwasm_std::AnyMsg {
77                    type_url: Self::TYPE_URL.to_string(),
78                    value: cosmwasm_std::Binary::from(self.to_proto_bytes()),
79                }
80            }
81        }
82
83        #query_request_conversion
84
85        impl From<#ident> for cosmwasm_std::Binary {
86            fn from(msg: #ident) -> Self {
87                cosmwasm_std::Binary::new(msg.to_proto_bytes())
88            }
89        }
90
91        impl TryFrom<cosmwasm_std::Binary> for #ident {
92            type Error = cosmwasm_std::StdError;
93
94            fn try_from(binary: cosmwasm_std::Binary) -> ::std::result::Result<Self, Self::Error> {
95                use ::prost::Message;
96                Ok(Self::decode(&binary[..]).map_err(|e| {
97                    cosmwasm_std::StdError::parse_err(
98                        stringify!(#ident),
99                        format!(
100                            "Unable to decode binary: \n  - base64: {}\n  - bytes array: {:?}\n\n{:?}",
101                            binary,
102                            binary.to_vec(),
103                            e
104                        )
105                    )
106                }).unwrap())
107            }
108        }
109
110        impl TryFrom<cosmwasm_std::SubMsgResult> for #ident {
111            type Error = cosmwasm_std::StdError;
112
113            fn try_from(result: cosmwasm_std::SubMsgResult) -> ::std::result::Result<Self, Self::Error> {
114                result
115                    .into_result()
116                    .map_err(|e| cosmwasm_std::StdError::generic_err(e))?
117                    .data
118                    .ok_or_else(|| cosmwasm_std::StdError::not_found("cosmwasm_std::SubMsgResult::<T>"))?
119                    .try_into()
120            }
121        }
122    })
123    .into()
124}
125
126fn get_type_url(attrs: &[syn::Attribute]) -> proc_macro2::TokenStream {
127    let proto_message = get_attr("proto_message", attrs).and_then(|a| a.parse_meta().ok());
128
129    if let Some(syn::Meta::List(meta)) = proto_message.clone() {
130        match meta.nested[0].clone() {
131            syn::NestedMeta::Meta(syn::Meta::NameValue(meta)) => {
132                if meta.path.is_ident("type_url") {
133                    match meta.lit {
134                        syn::Lit::Str(s) => quote!(#s),
135                        _ => proto_message_attr_error(meta.lit),
136                    }
137                } else {
138                    proto_message_attr_error(meta.path)
139                }
140            }
141            t => proto_message_attr_error(t),
142        }
143    } else {
144        proto_message_attr_error(proto_message)
145    }
146}
147
148fn get_query_attrs<F>(attrs: &[syn::Attribute], f: F) -> proc_macro2::TokenStream
149where
150    F: FnMut(&Vec<TokenTree>) -> Option<proc_macro2::TokenStream>,
151{
152    let proto_query = get_attr("proto_query", attrs);
153
154    if let Some(attr) = proto_query {
155        if attr.tokens.clone().into_iter().count() != 1 {
156            return proto_query_attr_error(proto_query);
157        }
158
159        if let Some(TokenTree::Group(group)) = attr.tokens.clone().into_iter().next() {
160            let kv_groups = group.stream().into_iter().group_by(|t| {
161                if let TokenTree::Punct(punct) = t {
162                    punct.as_char() != ','
163                } else {
164                    true
165                }
166            });
167            let mut key_values: Vec<Vec<TokenTree>> = vec![];
168
169            for (non_sep, g) in &kv_groups {
170                if non_sep {
171                    key_values.push(g.collect());
172                }
173            }
174
175            return key_values
176                .iter()
177                .find_map(f)
178                .unwrap_or_else(|| proto_query_attr_error(proto_query));
179        }
180
181        proto_query_attr_error(proto_query)
182    } else {
183        proto_query_attr_error(proto_query)
184    }
185}
186
187fn get_attr<'a>(attr_ident: &str, attrs: &'a [syn::Attribute]) -> Option<&'a syn::Attribute> {
188    attrs
189        .iter()
190        .find(|&attr| attr.path.segments.len() == 1 && attr.path.segments[0].ident == attr_ident)
191}
192
193fn proto_message_attr_error<T: quote::ToTokens>(tokens: T) -> proc_macro2::TokenStream {
194    syn::Error::new_spanned(tokens, "expected `proto_message(type_url = \"...\")`")
195        .to_compile_error()
196}
197
198fn proto_query_attr_error<T: quote::ToTokens>(tokens: T) -> proc_macro2::TokenStream {
199    syn::Error::new_spanned(
200        tokens,
201        "expected `proto_query(path = \"...\", response_type = ...)`",
202    )
203    .to_compile_error()
204}