Skip to main content

safe_math_macros/
lib.rs

1#![forbid(unsafe_code)]
2
3use proc_macro::TokenStream;
4use quote::{format_ident, quote};
5use std::sync::atomic::{AtomicUsize, Ordering};
6use syn::{
7    fold::{self, Fold},
8    parse_macro_input,
9    spanned::Spanned,
10    BinOp, Expr, ExprBinary, ItemFn,
11};
12#[cfg(feature = "derive")]
13mod derive;
14
15// Global counter for generating unique variable names
16static TEMP_VAR_COUNTER: AtomicUsize = AtomicUsize::new(0);
17
18#[proc_macro_attribute]
19pub fn safe_math(_attr: TokenStream, item: TokenStream) -> TokenStream {
20    let mut input_fn = parse_macro_input!(item as ItemFn);
21    let orig_block = *input_fn.block;
22
23    // ensure that the fn has a return type
24    let return_type = match &input_fn.sig.output {
25        syn::ReturnType::Type(_, ty) => ty,
26        syn::ReturnType::Default => {
27            return syn::Error::new(input_fn.sig.output.span(), "Function must return a Result")
28                .to_compile_error()
29                .into();
30        }
31    };
32
33    // ensure that the return type is a Result
34    let is_result = match &**return_type {
35        syn::Type::Path(type_path) => {
36            let segments = &type_path.path.segments;
37            segments
38                .last()
39                .map(|seg| seg.ident == "Result")
40                .unwrap_or(false)
41        }
42        _ => false,
43    };
44
45    if !is_result {
46        return syn::Error::new(return_type.span(), "Function must return a Result")
47            .to_compile_error()
48            .into();
49    }
50
51    let new_block = MathRewriter.fold_block(orig_block);
52    input_fn.block = Box::new(new_block);
53    TokenStream::from(quote! { #input_fn })
54}
55
56#[proc_macro]
57pub fn safe_math_block(input: TokenStream) -> TokenStream {
58    let expression = parse_macro_input!(input as syn::Expr);
59    let rewritten_expr = MathRewriter.fold_expr(expression);
60    TokenStream::from(quote! { #rewritten_expr })
61}
62
63/// Generates a unique variable name that is extremely unlikely to collide
64/// with user-defined variables
65fn generate_unique_temp_var() -> syn::Ident {
66    let counter = TEMP_VAR_COUNTER.fetch_add(1, Ordering::SeqCst);
67    // Use a very distinctive prefix that users are unlikely to use
68    // Include the counter to ensure uniqueness across multiple macro invocations
69    format_ident!(
70        "__safe_math_temp_ref_{}_{}",
71        std::process::id(), // Process ID for uniqueness across processes
72        counter             // Counter for uniqueness within process
73    )
74}
75
76pub(crate) struct MathRewriter;
77
78impl Fold for MathRewriter {
79    fn fold_expr(&mut self, expr: Expr) -> Expr {
80        match expr {
81            Expr::Binary(ExprBinary {
82                left,
83                op: BinOp::Add(_),
84                right,
85                ..
86            }) => {
87                let left = self.fold_expr(*left);
88                let right = self.fold_expr(*right);
89                syn::parse_quote! { ::safe_math::safe_add(#left, #right)? }
90            }
91            Expr::Binary(ExprBinary {
92                left,
93                op: BinOp::Sub(_),
94                right,
95                ..
96            }) => {
97                let left = self.fold_expr(*left);
98                let right = self.fold_expr(*right);
99                syn::parse_quote! { ::safe_math::safe_sub(#left, #right)? }
100            }
101            Expr::Binary(ExprBinary {
102                left,
103                op: BinOp::Mul(_),
104                right,
105                ..
106            }) => {
107                let left = self.fold_expr(*left);
108                let right = self.fold_expr(*right);
109                syn::parse_quote! { ::safe_math::safe_mul(#left, #right)? }
110            }
111            Expr::Binary(ExprBinary {
112                left,
113                op: BinOp::Div(_),
114                right,
115                ..
116            }) => {
117                let left = self.fold_expr(*left);
118                let right = self.fold_expr(*right);
119                syn::parse_quote! { ::safe_math::safe_div(#left, #right)? }
120            }
121            Expr::Binary(ExprBinary {
122                left,
123                op: BinOp::Rem(_),
124                right,
125                ..
126            }) => {
127                let left = self.fold_expr(*left);
128                let right = self.fold_expr(*right);
129                syn::parse_quote! { ::safe_math::safe_rem(#left, #right)? }
130            }
131            // Handle compound assignments by transforming them to regular assignments
132            // to avoid double evaluation of the left-hand side
133            Expr::Binary(ExprBinary {
134                left,
135                op: BinOp::AddAssign(_),
136                right,
137                ..
138            }) => {
139                let right = self.fold_expr(*right);
140                let temp_var = generate_unique_temp_var();
141                syn::parse_quote! {
142                    {
143                        let #temp_var = &mut #left;
144                        *#temp_var = ::safe_math::safe_add(*#temp_var, #right)?;
145                    }
146                }
147            }
148            Expr::Binary(ExprBinary {
149                left,
150                op: BinOp::SubAssign(_),
151                right,
152                ..
153            }) => {
154                let right = self.fold_expr(*right);
155                let temp_var = generate_unique_temp_var();
156                syn::parse_quote! {
157                    {
158                        let #temp_var = &mut #left;
159                        *#temp_var = ::safe_math::safe_sub(*#temp_var, #right)?;
160                    }
161                }
162            }
163            Expr::Binary(ExprBinary {
164                left,
165                op: BinOp::MulAssign(_),
166                right,
167                ..
168            }) => {
169                let right = self.fold_expr(*right);
170                let temp_var = generate_unique_temp_var();
171                syn::parse_quote! {
172                    {
173                        let #temp_var = &mut #left;
174                        *#temp_var = ::safe_math::safe_mul(*#temp_var, #right)?;
175                    }
176                }
177            }
178            Expr::Binary(ExprBinary {
179                left,
180                op: BinOp::DivAssign(_),
181                right,
182                ..
183            }) => {
184                let right = self.fold_expr(*right);
185                let temp_var = generate_unique_temp_var();
186                syn::parse_quote! {
187                    {
188                        let #temp_var = &mut #left;
189                        *#temp_var = ::safe_math::safe_div(*#temp_var, #right)?;
190                    }
191                }
192            }
193            Expr::Binary(ExprBinary {
194                left,
195                op: BinOp::RemAssign(_),
196                right,
197                ..
198            }) => {
199                let right = self.fold_expr(*right);
200                let temp_var = generate_unique_temp_var();
201                syn::parse_quote! {
202                    {
203                        let #temp_var = &mut #left;
204                        *#temp_var = ::safe_math::safe_rem(*#temp_var, #right)?;
205                    }
206                }
207            }
208            _ => fold::fold_expr(self, expr),
209        }
210    }
211}
212
213#[cfg(feature = "derive")]
214#[proc_macro_derive(SafeMathOps, attributes(SafeMathOps))]
215pub fn derive_safe_math_ops(input: TokenStream) -> TokenStream {
216    derive::derive_safe_math_ops(input)
217}