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
15static 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 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 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
63fn generate_unique_temp_var() -> syn::Ident {
66 let counter = TEMP_VAR_COUNTER.fetch_add(1, Ordering::SeqCst);
67 format_ident!(
70 "__safe_math_temp_ref_{}_{}",
71 std::process::id(), counter )
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 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}