Skip to main content

sui_proc_macros/
lib.rs

1// Copyright (c) Mysten Labs, Inc.
2// SPDX-License-Identifier: Apache-2.0
3
4use proc_macro::TokenStream;
5use quote::{ToTokens, quote, quote_spanned};
6use syn::{
7    Attribute, BinOp, Data, DataEnum, DeriveInput, Expr, ExprBinary, ExprMacro, Item, ItemMacro,
8    Stmt, StmtMacro, Token, UnOp,
9    fold::{Fold, fold_expr, fold_item_macro, fold_stmt},
10    parse::Parser,
11    parse_macro_input, parse2,
12    punctuated::Punctuated,
13    spanned::Spanned,
14};
15
16/// The sui_test macro will invoke either `#[msim::test]` or `#[tokio::test]`,
17/// depending on whether the simulator config var is enabled.
18///
19/// This should be used for tests that can meaningfully run in either environment.
20#[proc_macro_attribute]
21pub fn sui_test(args: TokenStream, item: TokenStream) -> TokenStream {
22    let input = parse_macro_input!(item as syn::ItemFn);
23    let arg_parser = Punctuated::<syn::Meta, Token![,]>::parse_terminated;
24    let args = arg_parser.parse(args).unwrap().into_iter();
25
26    let header = if cfg!(msim) {
27        quote! {
28            #[::sui_simulator::sim_test(crate = "sui_simulator", #(#args)* )]
29        }
30    } else {
31        quote! {
32            #[::tokio::test(#(#args)*)]
33        }
34    };
35
36    let result = quote! {
37        #header
38        #input
39    };
40
41    result.into()
42}
43
44/// The sim_test macro will invoke `#[msim::test]` if the simulator config var is enabled.
45///
46/// Otherwise, it will emit an ignored test - if forcibly run, the ignored test will panic.
47///
48/// This macro must be used in order to pass any simulator-specific arguments (e.g. a
49/// custom `config`), which are not understood by tokio.
50#[proc_macro_attribute]
51pub fn sim_test(args: TokenStream, item: TokenStream) -> TokenStream {
52    let input = parse_macro_input!(item as syn::ItemFn);
53    let arg_parser = Punctuated::<syn::Meta, Token![,]>::parse_terminated;
54    let args = arg_parser.parse(args).unwrap().into_iter();
55
56    let ignore = input
57        .attrs
58        .iter()
59        .find(|attr| attr.path().is_ident("ignore"))
60        .map_or(quote! {}, |_| quote! { #[ignore] });
61
62    let result = if cfg!(msim) {
63        let sig = &input.sig;
64        let return_type = &sig.output;
65        let body = &input.block;
66        quote! {
67            #[::sui_simulator::sim_test(crate = "sui_simulator", #(#args),*)]
68            #ignore
69            #sig {
70                async fn body_fn() #return_type { #body }
71
72                let timeout_secs: u64 = std::env::var("SUI_SIM_TEST_TIMEOUT_SECS")
73                    .ok()
74                    .and_then(|s| s.parse().ok())
75                    .unwrap_or(1000);
76                let timeout_duration = tokio::time::Duration::from_secs(timeout_secs);
77
78                let ret = tokio::time::timeout(timeout_duration, body_fn())
79                    .await
80                    .expect("sim_test timed out");
81
82                ::sui_simulator::task::shutdown_all_nodes();
83
84                // all node handles should have been dropped after the above block exits, but task
85                // shutdown is asynchronous, so we need a brief delay before checking for leaks.
86                tokio::time::sleep(tokio::time::Duration::from_millis(10)).await;
87
88                assert_eq!(
89                    sui_simulator::NodeLeakDetector::get_current_node_count(),
90                    0,
91                    "SuiNode leak detected"
92                );
93
94                ret
95            }
96        }
97    } else {
98        let fn_name = &input.sig.ident;
99        let sig = &input.sig;
100        let body = &input.block;
101        quote! {
102            #[allow(clippy::needless_return)]
103            #[tokio::test]
104            #ignore
105            #sig {
106                if std::env::var("SUI_SKIP_SIMTESTS").is_ok() {
107                    println!("not running test {} in `cargo test`: SUI_SKIP_SIMTESTS is set", stringify!(#fn_name));
108
109                    struct Ret;
110
111                    impl From<Ret> for () {
112                        fn from(_ret: Ret) -> Self {
113                        }
114                    }
115
116                    impl<E> From<Ret> for Result<(), E> {
117                        fn from(_ret: Ret) -> Self {
118                            Ok(())
119                        }
120                    }
121
122                    return Ret.into();
123                }
124
125                #body
126            }
127        }
128    };
129
130    result.into()
131}
132
133#[proc_macro]
134pub fn checked_arithmetic(input: TokenStream) -> TokenStream {
135    let input_file = CheckArithmetic.fold_file(parse_macro_input!(input));
136
137    let output_items = input_file.items;
138
139    let output = quote! {
140        #(#output_items)*
141    };
142
143    TokenStream::from(output)
144}
145
146#[proc_macro_attribute]
147pub fn with_checked_arithmetic(_attr: TokenStream, item: TokenStream) -> TokenStream {
148    let input_item = parse_macro_input!(item as Item);
149    match input_item {
150        Item::Fn(input_fn) => {
151            let transformed_fn = CheckArithmetic.fold_item_fn(input_fn);
152            TokenStream::from(quote! { #transformed_fn })
153        }
154        Item::Impl(input_impl) => {
155            let transformed_impl = CheckArithmetic.fold_item_impl(input_impl);
156            TokenStream::from(quote! { #transformed_impl })
157        }
158        item => {
159            let transformed_impl = CheckArithmetic.fold_item(item);
160            TokenStream::from(quote! { #transformed_impl })
161        }
162    }
163}
164
165struct CheckArithmetic;
166
167impl CheckArithmetic {
168    fn maybe_skip_macro(&self, attrs: &mut Vec<Attribute>) -> bool {
169        if let Some(idx) = attrs
170            .iter()
171            .position(|attr| attr.path().is_ident("skip_checked_arithmetic"))
172        {
173            // Skip processing macro because it is annotated with
174            // #[skip_checked_arithmetic]
175            attrs.remove(idx);
176            true
177        } else {
178            false
179        }
180    }
181
182    fn process_macro_contents(
183        &mut self,
184        tokens: proc_macro2::TokenStream,
185    ) -> syn::Result<proc_macro2::TokenStream> {
186        // Parse the macro's contents as a comma-separated list of expressions.
187        let parser = Punctuated::<Expr, Token![,]>::parse_terminated;
188        let Ok(exprs) = parser.parse(tokens.clone().into()) else {
189            return Err(syn::Error::new_spanned(
190                tokens,
191                "could not process macro contents - use #[skip_checked_arithmetic] to skip this macro",
192            ));
193        };
194
195        // Fold each sub expression.
196        let folded_exprs = exprs
197            .into_iter()
198            .map(|expr| self.fold_expr(expr))
199            .collect::<Vec<_>>();
200
201        // Convert the folded expressions back into tokens and reconstruct the macro.
202        let mut folded_tokens = proc_macro2::TokenStream::new();
203        for (i, folded_expr) in folded_exprs.into_iter().enumerate() {
204            if i > 0 {
205                folded_tokens.extend(std::iter::once::<proc_macro2::TokenTree>(
206                    proc_macro2::Punct::new(',', proc_macro2::Spacing::Alone).into(),
207                ));
208            }
209            folded_expr.to_tokens(&mut folded_tokens);
210        }
211
212        Ok(folded_tokens)
213    }
214}
215
216impl Fold for CheckArithmetic {
217    fn fold_stmt(&mut self, stmt: Stmt) -> Stmt {
218        let stmt = fold_stmt(self, stmt);
219        if let Stmt::Macro(stmt_macro) = stmt {
220            let StmtMacro {
221                mut attrs,
222                mut mac,
223                semi_token,
224            } = stmt_macro;
225
226            if self.maybe_skip_macro(&mut attrs) {
227                Stmt::Macro(StmtMacro {
228                    attrs,
229                    mac,
230                    semi_token,
231                })
232            } else {
233                match self.process_macro_contents(mac.tokens.clone()) {
234                    Ok(folded_tokens) => {
235                        mac.tokens = folded_tokens;
236                        Stmt::Macro(StmtMacro {
237                            attrs,
238                            mac,
239                            semi_token,
240                        })
241                    }
242                    Err(error) => parse2(error.to_compile_error()).unwrap(),
243                }
244            }
245        } else {
246            stmt
247        }
248    }
249
250    fn fold_item_macro(&mut self, mut item_macro: ItemMacro) -> ItemMacro {
251        if !self.maybe_skip_macro(&mut item_macro.attrs) {
252            let err = syn::Error::new_spanned(
253                item_macro.to_token_stream(),
254                "cannot process macros - use #[skip_checked_arithmetic] to skip \
255                    processing this macro",
256            );
257
258            return parse2(err.to_compile_error()).unwrap();
259        }
260        fold_item_macro(self, item_macro)
261    }
262
263    fn fold_expr(&mut self, expr: Expr) -> Expr {
264        let span = expr.span();
265        let expr = fold_expr(self, expr);
266        let expr = match expr {
267            Expr::Macro(expr_macro) => {
268                let ExprMacro { mut attrs, mut mac } = expr_macro;
269
270                if self.maybe_skip_macro(&mut attrs) {
271                    return Expr::Macro(ExprMacro { attrs, mac });
272                } else {
273                    match self.process_macro_contents(mac.tokens.clone()) {
274                        Ok(folded_tokens) => {
275                            mac.tokens = folded_tokens;
276                            let expr_macro = Expr::Macro(ExprMacro { attrs, mac });
277                            quote!(#expr_macro)
278                        }
279                        Err(error) => {
280                            return Expr::Verbatim(error.to_compile_error());
281                        }
282                    }
283                }
284            }
285
286            Expr::Binary(expr_binary) => {
287                let ExprBinary {
288                    attrs,
289                    mut left,
290                    op,
291                    mut right,
292                } = expr_binary;
293
294                fn remove_parens(expr: &mut Expr) {
295                    if let Expr::Paren(paren) = expr {
296                        // i don't even think rust allows this, but just in case
297                        assert!(paren.attrs.is_empty(), "TODO: attrs on parenthesized");
298                        *expr = *paren.expr.clone();
299                    }
300                }
301
302                macro_rules! wrap_op {
303                    ($left: expr, $right: expr, $method: ident, $span: expr) => {{
304                        // Remove parens from exprs since both sides get assigned to tmp variables.
305                        // otherwise we get lint errors
306                        remove_parens(&mut $left);
307                        remove_parens(&mut $right);
308
309                        quote_spanned!($span => {
310                            // assign in one stmt in case either #left or #right contains
311                            // references to `left` or `right` symbols.
312                            let (left, right) = (#left, #right);
313                            left.$method(right)
314                                .unwrap_or_else(||
315                                    panic!(
316                                        "Overflow or underflow in {} {} + {}",
317                                        stringify!($method),
318                                        left,
319                                        right,
320                                    )
321                                )
322                        })
323                    }};
324                }
325
326                macro_rules! wrap_op_assign {
327                    ($left: expr, $right: expr, $method: ident, $span: expr) => {{
328                        // Remove parens from exprs since both sides get assigned to tmp variables.
329                        // otherwise we get lint errors
330                        remove_parens(&mut $left);
331                        remove_parens(&mut $right);
332
333                        quote_spanned!($span => {
334                            // assign in one stmt in case either #left or #right contains
335                            // references to `left` or `right` symbols.
336                            let (left, right) = (&mut #left, #right);
337                            *left = (*left).$method(right)
338                                .unwrap_or_else(||
339                                    panic!(
340                                        "Overflow or underflow in {} {} + {}",
341                                        stringify!($method),
342                                        *left,
343                                        right
344                                    )
345                                )
346                        })
347                    }};
348                }
349
350                match op {
351                    BinOp::Add(_) => {
352                        wrap_op!(left, right, checked_add, span)
353                    }
354                    BinOp::Sub(_) => {
355                        wrap_op!(left, right, checked_sub, span)
356                    }
357                    BinOp::Mul(_) => {
358                        wrap_op!(left, right, checked_mul, span)
359                    }
360                    BinOp::Div(_) => {
361                        wrap_op!(left, right, checked_div, span)
362                    }
363                    BinOp::Rem(_) => {
364                        wrap_op!(left, right, checked_rem, span)
365                    }
366                    BinOp::AddAssign(_) => {
367                        wrap_op_assign!(left, right, checked_add, span)
368                    }
369                    BinOp::SubAssign(_) => {
370                        wrap_op_assign!(left, right, checked_sub, span)
371                    }
372                    BinOp::MulAssign(_) => {
373                        wrap_op_assign!(left, right, checked_mul, span)
374                    }
375                    BinOp::DivAssign(_) => {
376                        wrap_op_assign!(left, right, checked_div, span)
377                    }
378                    BinOp::RemAssign(_) => {
379                        wrap_op_assign!(left, right, checked_rem, span)
380                    }
381                    _ => {
382                        let expr_binary = ExprBinary {
383                            attrs,
384                            left,
385                            op,
386                            right,
387                        };
388                        quote_spanned!(span => #expr_binary)
389                    }
390                }
391            }
392            Expr::Unary(expr_unary) => {
393                let op = &expr_unary.op;
394                let operand = &expr_unary.expr;
395                match op {
396                    UnOp::Neg(_) => {
397                        quote_spanned!(span => #operand.checked_neg().expect("Overflow or underflow in negation"))
398                    }
399                    _ => quote_spanned!(span => #expr_unary),
400                }
401            }
402            _ => quote_spanned!(span => #expr),
403        };
404
405        parse2(expr).unwrap()
406    }
407}
408
409/// This proc macro generates a function `order_to_variant_map` which returns a map
410/// of the position of each variant to the name of the variant.
411/// It is intended to catch changes in enum order when backward compat is required.
412/// ```rust,ignore
413///    /// Example for this enum
414///    #[derive(EnumVariantOrder)]
415///    pub enum MyEnum {
416///         A,
417///         B(u64),
418///         C{x: bool, y: i8},
419///     }
420///     let order_map = MyEnum::order_to_variant_map();
421///     assert!(order_map.get(0).unwrap() == "A");
422///     assert!(order_map.get(1).unwrap() == "B");
423///     assert!(order_map.get(2).unwrap() == "C");
424/// ```
425#[proc_macro_derive(EnumVariantOrder)]
426pub fn enum_variant_order_derive(input: TokenStream) -> TokenStream {
427    let ast = parse_macro_input!(input as DeriveInput);
428    let name = &ast.ident;
429
430    if let Data::Enum(DataEnum { variants, .. }) = ast.data {
431        let variant_entries = variants
432            .iter()
433            .enumerate()
434            .map(|(index, variant)| {
435                let variant_name = variant.ident.to_string();
436                quote! {
437                    map.insert( #index as u64, (#variant_name).to_string());
438                }
439            })
440            .collect::<Vec<_>>();
441
442        let deriv = quote! {
443            impl sui_enum_compat_util::EnumOrderMap for #name {
444                fn order_to_variant_map() -> std::collections::BTreeMap<u64, String > {
445                    let mut map = std::collections::BTreeMap::new();
446                    #(#variant_entries)*
447                    map
448                }
449            }
450        };
451
452        deriv.into()
453    } else {
454        panic!("EnumVariantOrder can only be used with enums.");
455    }
456}