1use 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#[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#[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 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 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 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 let folded_exprs = exprs
197 .into_iter()
198 .map(|expr| self.fold_expr(expr))
199 .collect::<Vec<_>>();
200
201 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 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(&mut $left);
307 remove_parens(&mut $right);
308
309 quote_spanned!($span => {
310 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(&mut $left);
331 remove_parens(&mut $right);
332
333 quote_spanned!($span => {
334 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#[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}