injective_std_derive/
lib.rs1use 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 let type_url = get_type_url(&input.attrs);
30
31 let (query_request_conversion, cosmwasm_query) = if get_attr("proto_query", &input.attrs).is_some() {
35 let path = get_query_attrs(&input.attrs, match_kv_attr!("path", Literal));
36 let res = get_query_attrs(&input.attrs, match_kv_attr!("response_type", Ident));
37
38 let query_request_conversion = quote! {
39 impl <Q: cosmwasm_std::CustomQuery> From<#ident> for cosmwasm_std::QueryRequest<Q> {
40 fn from(msg: #ident) -> Self {
41 cosmwasm_std::QueryRequest::<Q>::Stargate {
42 path: #path.to_string(),
43 data: msg.into(),
44 }
45 }
46 }
47 };
48
49 let cosmwasm_query = quote! {
50 pub fn query(self, querier: &cosmwasm_std::QuerierWrapper<impl cosmwasm_std::CustomQuery>) -> cosmwasm_std::StdResult<#res> {
51 querier.query::<#res>(&self.into())
52 }
53 };
54 (query_request_conversion, cosmwasm_query)
55 } else {
56 (quote!(), quote!())
57 };
58
59 (quote! {
60 impl #ident {
61 pub const TYPE_URL: &'static str = #type_url;
62 #cosmwasm_query
63 }
64
65 #query_request_conversion
66
67 impl From<#ident> for cosmwasm_std::Binary {
68 fn from(msg: #ident) -> Self {
69 let mut bytes = Vec::new();
70 prost::Message::encode(&msg, &mut bytes)
71 .expect("Message encoding must be infallible");
72 cosmwasm_std::Binary::new(bytes)
73 }
74 }
75
76 impl<T> From<#ident> for cosmwasm_std::CosmosMsg<T> {
77 fn from(msg: #ident) -> Self {
78 cosmwasm_std::CosmosMsg::<T>::Stargate {
79 type_url: #type_url.to_string(),
80 value: msg.into(),
81 }
82 }
83 }
84
85 impl TryFrom<cosmwasm_std::Binary> for #ident {
86 type Error = cosmwasm_std::StdError;
87 fn try_from(binary: cosmwasm_std::Binary) -> Result<Self, Self::Error> {
88 use ::prost::Message;
89 Self::decode(&binary[..]).map_err(|e| {
90 cosmwasm_std::StdError::msg(format!(
91 "Unable to decode binary for {}: base64: {}, bytes: {:?}, error: {:?}",
92 stringify!(#ident),
93 binary,
94 binary.to_vec(),
95 e
96 ))
97 })
98 }
99 }
100
101 impl TryFrom<cosmwasm_std::SubMsgResult> for #ident {
102 type Error = cosmwasm_std::StdError;
103 fn try_from(result: cosmwasm_std::SubMsgResult) -> Result<Self, Self::Error> {
104 result
105 .into_result()
106 .map_err(|e| cosmwasm_std::StdError::msg(format!("SubMsgResult error: {}", e)))?
107 .data
108 .ok_or_else(|| {
109 cosmwasm_std::StdError::msg("No data found in SubMsgResult".to_string())
110 })?
111 .try_into()
112 }
113 }
114 })
115 .into()
116}
117
118fn get_type_url(attrs: &[syn::Attribute]) -> proc_macro2::TokenStream {
119 let proto_message = get_attr("proto_message", attrs).and_then(|a| a.parse_meta().ok());
120 if let Some(syn::Meta::List(meta)) = proto_message.clone() {
121 match meta.nested[0].clone() {
122 syn::NestedMeta::Meta(syn::Meta::NameValue(meta)) => {
123 if meta.path.is_ident("type_url") {
124 match meta.lit {
125 syn::Lit::Str(s) => quote!(#s),
126 _ => proto_message_attr_error(meta.lit),
127 }
128 } else {
129 proto_message_attr_error(meta.path)
130 }
131 }
132 t => proto_message_attr_error(t),
133 }
134 } else {
135 proto_message_attr_error(proto_message)
136 }
137}
138
139fn get_query_attrs<F>(attrs: &[syn::Attribute], f: F) -> proc_macro2::TokenStream
140where
141 F: FnMut(&Vec<TokenTree>) -> Option<proc_macro2::TokenStream>,
142{
143 let proto_query = get_attr("proto_query", attrs);
144 if let Some(attr) = proto_query {
145 if attr.tokens.clone().into_iter().count() != 1 {
146 return proto_query_attr_error(proto_query);
147 }
148 if let Some(TokenTree::Group(group)) = attr.tokens.clone().into_iter().next() {
149 let kv_groups = group
150 .stream()
151 .into_iter()
152 .chunk_by(|t| if let TokenTree::Punct(punct) = t { punct.as_char() != ',' } else { true });
153 let mut key_values: Vec<Vec<TokenTree>> = vec![];
154 for (non_sep, g) in &kv_groups {
155 if non_sep {
156 key_values.push(g.collect());
157 }
158 }
159 return key_values.iter().find_map(f).unwrap_or_else(|| proto_query_attr_error(proto_query));
160 }
161 proto_query_attr_error(proto_query)
162 } else {
163 proto_query_attr_error(proto_query)
164 }
165}
166
167fn get_attr<'a>(attr_ident: &str, attrs: &'a [syn::Attribute]) -> Option<&'a syn::Attribute> {
168 attrs
169 .iter()
170 .find(|&attr| attr.path.segments.len() == 1 && attr.path.segments[0].ident == attr_ident)
171}
172
173fn proto_message_attr_error<T: quote::ToTokens>(tokens: T) -> proc_macro2::TokenStream {
174 syn::Error::new_spanned(tokens, "expected `proto_message(type_url = \"...\")`").to_compile_error()
175}
176
177fn proto_query_attr_error<T: quote::ToTokens>(tokens: T) -> proc_macro2::TokenStream {
178 syn::Error::new_spanned(tokens, "expected `proto_query(path = \"...\", response_type = ...)`").to_compile_error()
179}