1use super::UnsafeOzErc721Mint;
2use crate::{
3 linter::{LateLintPass, LintContext},
4 sol::{
5 Severity, SolLint,
6 analysis::{
7 OPENZEPPELIN_ROOTS, arg_for_param, for_each_lhs_var, is_address_type, is_builtin,
8 is_literal_false, is_require_or_assert, loop_stmts, source_in_package, underlying_var,
9 unique, write_target,
10 },
11 },
12};
13use alloy_primitives::U256;
14use solar::{
15 ast::{ElementaryType, LitKind, StateMutability, Visibility},
16 interface::{Span, kw},
17 sema::{
18 Gcx,
19 hir::{
20 self, BinOpKind, CallArgs, Expr, ExprKind, FunctionId, Hir, ItemId, Stmt, StmtKind,
21 TypeKind, VariableId, Visit,
22 },
23 ty::{TyFn, TyKind},
24 },
25};
26use std::{ops::ControlFlow, slice};
27
28declare_forge_lint!(
29 UNSAFE_OZ_ERC721_MINT,
30 Severity::Med,
31 "unsafe-oz-erc721-mint",
32 "`ERC721._mint` does not check that the recipient can receive the token; use `_safeMint`"
33);
34
35impl<'gcx> LateLintPass<'gcx> for UnsafeOzErc721Mint {
36 fn check_function(
37 &mut self,
38 ctx: &LintContext,
39 gcx: Gcx<'gcx>,
40 func: &'gcx hir::Function<'gcx>,
41 ) {
42 let cx = Cx { gcx };
43 if named(func, "_safeMint")
47 && func
48 .contract
49 .is_some_and(|id| is_canonical_erc721(gcx.hir.contract(id).name.as_str()))
50 && source_in_package(&gcx.hir, func.source, OPENZEPPELIN_ROOTS)
51 {
52 return;
53 }
54 if (named(func, "_mint") && func.override_) || cx.is_override_delegation_helper(func) {
60 return;
61 }
62 let Some(body) = &func.body else { return };
68 for (callee, _, span) in cx.calls(body.stmts) {
69 let helper = cx.is_override_delegation_helper(gcx.hir.function(callee));
70 if cx.unsafe_mint_target(callee, helper, &mut Vec::new()).is_some() {
71 ctx.emit(&UNSAFE_OZ_ERC721_MINT, span);
72 }
73 }
74 }
75}
76
77#[derive(Clone, Copy)]
81struct UnsafeMintTarget {
82 preserves_recipient: bool,
83 preserves_token: bool,
84 preserves_code_length: bool,
85}
86
87type Call<'gcx> = (FunctionId, &'gcx CallArgs<'gcx>, Span);
89
90#[derive(Clone, Copy)]
92struct Cx<'gcx> {
93 gcx: Gcx<'gcx>,
94}
95
96impl<'gcx> Cx<'gcx> {
97 fn is_override_delegation_helper(self, function: &'gcx hir::Function<'gcx>) -> bool {
100 if !is_internal(function) || (function.override_ && named(function, "_mint")) {
101 return false;
102 }
103 let Some(contract_id) = function.contract else { return false };
104 let Some(function_id) = self
105 .gcx
106 .hir
107 .contract(contract_id)
108 .all_functions()
109 .find(|&id| std::ptr::eq(self.gcx.hir.function(id), function))
110 else {
111 return false;
112 };
113 self.gcx.hir.contract_ids().any(|candidate| {
114 let candidate = self.gcx.hir.contract(candidate);
115 candidate.linearized_bases.contains(&contract_id)
116 && candidate.all_functions().any(|id| {
117 let f = self.gcx.hir.function(id);
118 f.override_
119 && named(f, "_mint")
120 && self.function_reaches(id, function_id, &mut Vec::new())
121 })
122 })
123 }
124
125 fn function_reaches(
127 self,
128 function_id: FunctionId,
129 target: FunctionId,
130 seen: &mut Vec<FunctionId>,
131 ) -> bool {
132 if seen.contains(&function_id) {
133 return false;
134 }
135 seen.push(function_id);
136 let Some(body) = self.gcx.hir.function(function_id).body else { return false };
137 self.calls(body.stmts).iter().any(|&(callee, ..)| {
138 callee == target
139 || (is_internal(self.gcx.hir.function(callee))
140 && self.function_reaches(callee, target, seen))
141 })
142 }
143
144 fn unsafe_mint_target(
152 self,
153 function_id: FunctionId,
154 helper: bool,
155 seen: &mut Vec<FunctionId>,
156 ) -> Option<UnsafeMintTarget> {
157 if seen.contains(&function_id) {
158 return None;
159 }
160 seen.push(function_id);
161 let function = self.gcx.hir.function(function_id);
162 let is_mint = named(function, "_mint");
163 if !(is_mint || (helper && is_internal(function))) {
164 return None;
165 }
166 let contract = self.gcx.hir.contract(function.contract?);
167 if contract.kind.is_library() {
168 return None;
169 }
170 let canonical = is_canonical_erc721(contract.name.as_str())
171 && source_in_package(&self.gcx.hir, function.source, OPENZEPPELIN_ROOTS);
172 if canonical && named(function, "_safeMint") {
173 return None;
174 }
175 if canonical && is_mint {
178 return Some(UnsafeMintTarget {
179 preserves_recipient: true,
180 preserves_token: true,
181 preserves_code_length: true,
182 });
183 }
184 if !(function.override_ || helper) {
185 return None;
186 }
187 let body = function.body.as_ref()?;
188 let recipient =
190 function.parameters.iter().copied().find(|&vid| is_address_type(&self.gcx.hir, vid));
191 let calls = self.calls(body.stmts);
192 let mut unsafe_targets = Vec::new();
196 let mut unstable_code_targets = Vec::new();
197 let mut judged = Vec::new();
198 let (mut targets_preserve_recipient, mut targets_preserve_token) = (true, true);
199 for &(callee, ..) in &calls {
200 if judged.contains(&callee) {
201 continue;
202 }
203 judged.push(callee);
204 if let Some(target) = self.unsafe_mint_target(callee, true, &mut seen.clone()) {
205 unsafe_targets.push(callee);
206 if !target.preserves_code_length {
207 unstable_code_targets.push(callee);
208 }
209 targets_preserve_recipient &= target.preserves_recipient;
210 targets_preserve_token &= target.preserves_token;
211 }
212 }
213 let delegations: Vec<_> =
214 calls.iter().filter(|(callee, ..)| unsafe_targets.contains(callee)).collect();
215 if delegations.is_empty() {
216 return None;
217 }
218 let forwards = |index: usize, var: Option<VariableId>| {
222 var.is_some_and(|var| {
223 delegations.iter().all(|&&(callee, args, _)| {
224 self.arg(callee, args, index).and_then(|expr| underlying_var(self.gcx, expr))
225 == Some(var)
226 })
227 })
228 };
229 let only_to_recipient = forwards(0, recipient);
230 let mut token = None;
235 let mut token_consistent = true;
236 for &&(callee, args, _) in &delegations {
237 let minted = self.arg(callee, args, 1).and_then(|expr| underlying_var(self.gcx, expr));
238 match minted.filter(|&minted| keeps_its_value(self.gcx, minted)) {
239 Some(minted) => {
240 token_consistent &= token.is_none_or(|token| token == minted);
241 token = Some(minted);
242 }
243 None => token_consistent = false,
244 }
245 }
246 let guarded = |recipient, token, seed| {
250 let mut walk = self.modifier_coverage_at_body(function, recipient, token, seed);
251 let mut walker = GuardWalker {
252 cx: self,
253 recipient,
254 token,
255 delegations: &unsafe_targets,
256 unstable_code_delegations: &unstable_code_targets,
257 seen: &mut Vec::new(),
258 };
259 walker.walk(body.stmts, &mut walk);
260 !walk.failed && !walk.pending
261 };
262 if only_to_recipient
267 && targets_preserve_recipient
268 && let Some(recipient) = recipient
269 && guarded(recipient, recipient, GuardCoverage::None)
270 {
271 return None;
272 }
273 if only_to_recipient
274 && token_consistent
275 && targets_preserve_recipient
276 && targets_preserve_token
277 && let Some(recipient) = recipient
278 && let Some(token) = token
279 && guarded(recipient, token, GuardCoverage::None)
280 {
281 return None;
282 }
283 let preserves = |index: usize| {
288 function.parameters.get(index).is_some_and(|&var| {
289 !body.stmts.iter().any(|stmt| self.mutates_var(stmt, var))
290 && !function.modifiers.iter().any(|modifier| {
291 modifier.args.exprs().any(|arg| self.expr_mutates_var(arg, var))
292 })
293 && forwards(index, Some(var))
294 })
295 };
296 let preserves_code_length = recipient
299 .is_some_and(|recipient| guarded(recipient, recipient, GuardCoverage::CodeLess));
300 Some(UnsafeMintTarget {
301 preserves_recipient: targets_preserve_recipient && preserves(0),
302 preserves_token: targets_preserve_token && preserves(1),
303 preserves_code_length,
304 })
305 }
306
307 fn any_in_stmts(
309 self,
310 stmts: &'gcx [Stmt<'gcx>],
311 stmt_matches: impl FnMut(&'gcx Stmt<'gcx>) -> bool,
312 expr_matches: impl FnMut(&'gcx Expr<'gcx>) -> bool,
313 ) -> bool {
314 let mut finder = Finder { gcx: self.gcx, stmt_matches, expr_matches };
315 stmts.iter().any(|stmt| finder.visit_stmt(stmt).is_break())
316 }
317
318 fn any_in_expr(
319 self,
320 expr: &'gcx Expr<'gcx>,
321 expr_matches: impl FnMut(&'gcx Expr<'gcx>) -> bool,
322 ) -> bool {
323 Finder { gcx: self.gcx, stmt_matches: |_| false, expr_matches }.visit_expr(expr).is_break()
324 }
325
326 fn calls(self, stmts: &'gcx [Stmt<'gcx>]) -> Vec<Call<'gcx>> {
328 let mut calls = Vec::new();
329 self.any_in_stmts(
330 stmts,
331 |_| false,
332 |expr| {
333 if let ExprKind::Call(_, args) = &expr.kind
334 && let Some(function_id) = self.resolved_callee(expr)
335 {
336 calls.push((function_id, args, expr.span));
337 }
338 false
339 },
340 );
341 calls
342 }
343
344 fn resolved_callee(self, expr: &Expr<'_>) -> Option<FunctionId> {
346 let ExprKind::Call(callee, ..) = &expr.kind else { return None };
347 self.gcx.resolved_function(callee)
348 }
349
350 fn callee_fn(self, expr: &Expr<'_>) -> Option<&'gcx TyFn<'gcx>> {
351 let ExprKind::Call(callee, ..) = &expr.kind else { return None };
352 match self.gcx.type_of_expr(callee.peel_parens().id)?.kind {
353 TyKind::Fn(function_ty) => Some(function_ty),
354 _ => None,
355 }
356 }
357
358 fn resolved_internal_callee(self, expr: &Expr<'_>) -> Option<FunctionId> {
362 let function_ty = self.callee_fn(expr)?;
363 function_ty.is_internal().then_some(function_ty.function_id).flatten()
364 }
365
366 fn is_unresolved_internal_pointer_call(self, expr: &Expr<'_>) -> bool {
370 let ExprKind::Call(callee, ..) = &expr.kind else { return false };
371 self.callee_fn(expr).is_some_and(|f| f.is_internal() && f.function_id.is_none())
372 && matches!(callee.peel_parens().kind, ExprKind::Ident(_))
373 && self.gcx.resolved_variable(callee).is_some()
374 }
375
376 fn arg(
378 self,
379 function_id: FunctionId,
380 args: &'gcx CallArgs<'gcx>,
381 index: usize,
382 ) -> Option<&'gcx Expr<'gcx>> {
383 let function = self.gcx.hir.function(function_id);
384 arg_for_param(self.gcx, function_id, *function.parameters.get(index)?, args)
385 }
386
387 fn is_receiver_hook(self, function_id: FunctionId) -> bool {
393 let function = self.gcx.hir.function(function_id);
394 let Some(contract) = function.contract else { return false };
395 let &[from, to, id, data] = function.parameters else { return false };
396 let kind = |vid: VariableId| &self.gcx.hir.variable(vid).ty.kind;
397 named(function, "onERC721Received")
398 && !self.gcx.hir.contract(contract).kind.is_library()
399 && matches!(function.visibility, Visibility::Public | Visibility::External)
400 && is_address_type(&self.gcx.hir, from)
401 && is_address_type(&self.gcx.hir, to)
402 && matches!(kind(id), TypeKind::Elementary(ElementaryType::UInt(_)))
403 && matches!(kind(data), TypeKind::Elementary(ElementaryType::Bytes))
404 }
405
406 fn is_received_selector(self, expr: &Expr<'gcx>) -> bool {
412 let expr = expr.peel_parens();
413 match &expr.kind {
414 ExprKind::Lit(lit) => {
415 matches!(&lit.kind, LitKind::Number(value) if *value == U256::from(ERC721_RECEIVED))
416 }
417 ExprKind::Call(callee, args)
418 if matches!(callee.peel_parens().kind, ExprKind::Type(..)) =>
419 {
420 args.len() == 1
421 && args.exprs().next().is_some_and(|inner| {
422 self.selector_cast_preserves(expr, inner)
423 && self.is_received_selector(inner)
424 })
425 }
426 ExprKind::Member(base, member) => {
427 member.as_str() == "selector"
428 && self.gcx.resolved_function(base).is_some_and(|id| self.is_receiver_hook(id))
429 }
430 ExprKind::Ident(_) => self.gcx.resolved_variable(expr).is_some_and(|vid| {
432 let variable = self.gcx.hir.variable(vid);
433 variable.is_constant()
434 && variable.initializer.is_some_and(|init| self.is_received_selector(init))
435 }),
436 _ => false,
437 }
438 }
439
440 fn selector_cast_preserves(self, cast: &Expr<'_>, inner: &Expr<'_>) -> bool {
445 let encoding = |expr: &Expr<'_>| match self.gcx.type_of_expr(expr.peel_parens().id)?.kind {
446 TyKind::IntLiteral(..) => Some(SelectorEncoding::Literal),
447 TyKind::Elementary(ElementaryType::Int(size) | ElementaryType::UInt(size)) => {
448 Some(SelectorEncoding::Integer(size.bits()))
449 }
450 TyKind::Elementary(ElementaryType::FixedBytes(size)) => {
451 Some(SelectorEncoding::FixedBytes(size.bytes()))
452 }
453 _ => None,
454 };
455 matches!(
456 (encoding(inner), encoding(cast)),
457 (
458 Some(SelectorEncoding::Literal | SelectorEncoding::Integer(_)),
459 Some(SelectorEncoding::Integer(32..) | SelectorEncoding::FixedBytes(4))
460 ) | (Some(SelectorEncoding::FixedBytes(4)), Some(SelectorEncoding::Integer(32)))
461 | (
462 Some(SelectorEncoding::FixedBytes(4..)),
463 Some(SelectorEncoding::FixedBytes(4..))
464 )
465 )
466 }
467
468 fn branch_always_reverts(self, stmt: &'gcx Stmt<'gcx>) -> bool {
472 match &stmt.kind {
473 StmtKind::Revert(_) => !self.may_return(stmt),
474 StmtKind::Expr(expr) => is_revert_call(self.gcx, expr) && !self.may_return(stmt),
475 StmtKind::Block(block) | StmtKind::UncheckedBlock(block) => block
478 .stmts
479 .iter()
480 .find_map(|stmt| {
481 self.branch_always_reverts(stmt)
482 .then_some(true)
483 .or_else(|| self.may_return(stmt).then_some(false))
484 })
485 .unwrap_or(false),
486 StmtKind::If(cond, then, Some(otherwise)) => {
487 !self.expr_contains_frame_ending_assembly(cond)
488 && self.branch_always_reverts(then)
489 && self.branch_always_reverts(otherwise)
490 }
491 _ => false,
492 }
493 }
494
495 fn may_return(self, stmt: &'gcx Stmt<'gcx>) -> bool {
499 self.contains_frame_ending_assembly(slice::from_ref(stmt), &mut Vec::new())
500 || match &stmt.kind {
501 StmtKind::Block(block) | StmtKind::UncheckedBlock(block) => {
502 block.stmts.iter().any(|stmt| self.may_return(stmt))
503 }
504 StmtKind::Loop(block, source) => {
505 loop_stmts(*block, *source).any(|stmt| self.may_return(stmt))
506 }
507 StmtKind::If(_, then, otherwise) => {
508 self.may_return(then) || otherwise.is_some_and(|stmt| self.may_return(stmt))
509 }
510 StmtKind::Return(_)
511 | StmtKind::AssemblyBlock(_)
512 | StmtKind::Try(_)
513 | StmtKind::Switch(_) => true,
514 _ => false,
515 }
516 }
517
518 fn contains_frame_ending_assembly(
523 self,
524 stmts: &'gcx [Stmt<'gcx>],
525 seen: &mut Vec<FunctionId>,
526 ) -> bool {
527 self.any_in_stmts(stmts, is_assembly, |expr| self.call_leaves_frame(expr, seen))
528 }
529
530 fn expr_contains_frame_ending_assembly(self, expr: &'gcx Expr<'gcx>) -> bool {
531 let mut seen = Vec::new();
532 self.any_in_expr(expr, |expr| self.call_leaves_frame(expr, &mut seen))
533 }
534
535 fn call_leaves_frame(self, expr: &Expr<'_>, seen: &mut Vec<FunctionId>) -> bool {
536 self.is_unresolved_internal_pointer_call(expr)
537 || self
538 .resolved_internal_callee(expr)
539 .is_some_and(|id| self.callable_contains_frame_ending_assembly(id, seen))
540 }
541
542 fn callable_contains_frame_ending_assembly(
545 self,
546 function_id: FunctionId,
547 seen: &mut Vec<FunctionId>,
548 ) -> bool {
549 if seen.contains(&function_id) {
550 return false;
551 }
552 seen.push(function_id);
553 let function = self.gcx.hir.function(function_id);
554 let in_modifiers = function.modifiers.iter().any(|modifier| {
555 matches!(modifier.id, ItemId::Function(id)
556 if self.callable_contains_frame_ending_assembly(id, seen))
557 });
558 let in_body = function
559 .body
560 .as_ref()
561 .is_some_and(|body| self.contains_frame_ending_assembly(body.stmts, seen));
562 seen.pop();
563 in_modifiers || in_body
564 }
565
566 fn mutates_var(self, stmt: &'gcx Stmt<'gcx>, var: VariableId) -> bool {
572 self.any_in_stmts(slice::from_ref(stmt), is_assembly, |expr| {
573 assigns_to(self.gcx, expr, var)
574 })
575 }
576
577 fn expr_mutates_var(self, expr: &'gcx Expr<'gcx>, var: VariableId) -> bool {
578 self.any_in_expr(expr, |expr| assigns_to(self.gcx, expr, var))
579 }
580
581 fn stmts_may_change_account_code(
584 self,
585 stmts: &'gcx [Stmt<'gcx>],
586 delegations: &[FunctionId],
587 unstable_code_delegations: &[FunctionId],
588 seen: &mut Vec<FunctionId>,
589 ) -> bool {
590 self.any_in_stmts(stmts, is_assembly, |expr| {
591 self.call_may_change_account_code(expr, delegations, unstable_code_delegations, seen)
592 })
593 }
594
595 fn expr_may_change_account_code(
596 self,
597 expr: &'gcx Expr<'gcx>,
598 delegations: &[FunctionId],
599 unstable_code_delegations: &[FunctionId],
600 seen: &mut Vec<FunctionId>,
601 ) -> bool {
602 self.any_in_expr(expr, |expr| {
603 self.call_may_change_account_code(expr, delegations, unstable_code_delegations, seen)
604 })
605 }
606
607 fn call_may_change_account_code(
613 self,
614 expr: &Expr<'_>,
615 delegations: &[FunctionId],
616 unstable_code_delegations: &[FunctionId],
617 seen: &mut Vec<FunctionId>,
618 ) -> bool {
619 let ExprKind::Call(callee, ..) = &expr.kind else { return false };
620 let (callee, _) = callee.split_call_options();
621 let resolved = self.resolved_callee(expr);
622 if resolved.is_some_and(|id| delegations.contains(&id))
623 && !resolved.is_some_and(|id| unstable_code_delegations.contains(&id))
624 {
625 return false;
626 }
627 if matches!(callee.kind, ExprKind::New(_)) {
628 return true;
629 }
630 if !self.callee_fn(expr).is_some_and(|f| {
631 matches!(f.state_mutability, StateMutability::NonPayable | StateMutability::Payable)
632 }) {
633 return false;
634 }
635 match self.resolved_internal_callee(expr).filter(|&id| !self.gcx.hir.function(id).virtual_)
636 {
637 Some(id) => self.callable_may_change_account_code(id, seen),
638 None => true,
639 }
640 }
641
642 fn callable_may_change_account_code(
646 self,
647 function_id: FunctionId,
648 seen: &mut Vec<FunctionId>,
649 ) -> bool {
650 if seen.contains(&function_id) {
651 return false;
652 }
653 seen.push(function_id);
654 let function = self.gcx.hir.function(function_id);
655 let may_change = function.modifiers.iter().any(|modifier| {
656 modifier.args.exprs().any(|arg| self.expr_may_change_account_code(arg, &[], &[], seen))
657 }) || function.modifiers.iter().any(|modifier| {
658 matches!(modifier.id, ItemId::Function(id)
659 if self.callable_may_change_account_code(id, seen))
660 }) || function
661 .body
662 .as_ref()
663 .is_some_and(|body| self.stmts_may_change_account_code(body.stmts, &[], &[], seen));
664 seen.pop();
665 may_change
666 }
667
668 fn bound_guard_parameters(
670 self,
671 function_id: FunctionId,
672 args: &'gcx CallArgs<'gcx>,
673 recipient: VariableId,
674 token: VariableId,
675 ) -> Option<(VariableId, VariableId)> {
676 let parameters = self.gcx.hir.function(function_id).parameters;
677 let bound_to = |var| {
678 parameters
679 .iter()
680 .enumerate()
681 .find(|&(index, _)| {
682 self.arg(function_id, args, index)
683 .and_then(|expr| underlying_var(self.gcx, expr))
684 == Some(var)
685 })
686 .map(|(_, ¶meter)| parameter)
687 };
688 bound_to(recipient).zip(bound_to(token))
689 }
690
691 fn body_guards(
695 self,
696 function_id: FunctionId,
697 recipient: VariableId,
698 token: VariableId,
699 seen: &mut Vec<FunctionId>,
700 ) -> GuardCoverage {
701 if seen.contains(&function_id) {
702 return GuardCoverage::None;
703 }
704 seen.push(function_id);
705 let function = self.gcx.hir.function(function_id);
706 let guarded = match &function.body {
712 Some(body)
713 if !function.virtual_
714 && function.modifiers.is_empty()
715 && !body.stmts.iter().any(|stmt| {
716 self.mutates_var(stmt, recipient) || self.mutates_var(stmt, token)
717 }) =>
718 {
719 let mut walk = GuardWalk::default();
720 let mut walker = GuardWalker {
721 cx: self,
722 recipient,
723 token,
724 delegations: &[],
725 unstable_code_delegations: &[],
726 seen,
727 };
728 walker.walk(body.stmts, &mut walk);
729 if walk.escaped {
730 GuardCoverage::None
731 } else if walk.future_coverage == GuardCoverage::CodeLess {
732 GuardCoverage::CodeLess
733 } else if walk.coverage == GuardCoverage::CodeLess {
734 GuardCoverage::CallbackOrCodeLess
738 } else {
739 walk.coverage
740 }
741 }
742 _ => GuardCoverage::None,
743 };
744 seen.pop();
745 guarded
746 }
747
748 fn modifier_coverage_at_body(
754 self,
755 function: &'gcx hir::Function<'gcx>,
756 recipient: VariableId,
757 token: VariableId,
758 seed: GuardCoverage,
759 ) -> GuardWalk {
760 let mut state = GuardWalk { coverage: seed, future_coverage: seed, ..GuardWalk::default() };
761 let body_bypass = function
762 .body
763 .as_ref()
764 .is_some_and(|body| self.contains_frame_ending_assembly(body.stmts, &mut Vec::new()));
765 let mut has_tail_guard = false;
766 for (index, modifier) in function.modifiers.iter().enumerate() {
767 if modifier.args.exprs().any(|arg| {
768 self.expr_mutates_var(arg, recipient) || self.expr_mutates_var(arg, token)
769 }) {
770 state.coverage = GuardCoverage::None;
771 state.future_coverage = GuardCoverage::None;
772 has_tail_guard = false;
773 }
774 state.retire_code_snapshots_if(|| {
775 modifier
776 .args
777 .exprs()
778 .any(|arg| self.expr_may_change_account_code(arg, &[], &[], &mut Vec::new()))
779 });
780 let ItemId::Function(modifier_id) = modifier.id else { continue };
781 let Some(body) = &self.gcx.hir.function(modifier_id).body else { continue };
782 let Some((prefix, suffix)) = modifier_body_sides(body.stmts) else {
783 state.retire_code_snapshots_if(|| {
787 self.stmts_may_change_account_code(body.stmts, &[], &[], &mut Vec::new())
788 });
789 continue;
790 };
791 let prefix_may_change_code =
792 || self.stmts_may_change_account_code(prefix, &[], &[], &mut Vec::new());
793 let Some((modifier_recipient, modifier_token)) =
794 self.bound_guard_parameters(modifier_id, &modifier.args, recipient, token)
795 else {
796 state.retire_code_snapshots_if(prefix_may_change_code);
797 continue;
798 };
799 let parameters_unchanged = !body.stmts.iter().any(|stmt| {
800 self.mutates_var(stmt, modifier_recipient) || self.mutates_var(stmt, modifier_token)
801 });
802 let mut walker = GuardWalker {
803 cx: self,
804 recipient: modifier_recipient,
805 token: modifier_token,
806 delegations: &[],
807 unstable_code_delegations: &[],
808 seen: &mut Vec::new(),
809 };
810 if parameters_unchanged {
811 walker.walk(prefix, &mut state);
812 } else {
813 state.retire_code_snapshots_if(prefix_may_change_code);
814 }
815 let inner_modifier_bypass = function.modifiers[index + 1..].iter().any(|inner| {
816 matches!(inner.id, ItemId::Function(id)
817 if self.callable_contains_frame_ending_assembly(id, &mut Vec::new()))
818 });
819 if parameters_unchanged && !body_bypass && !inner_modifier_bypass {
820 let mut suffix_walk = GuardWalk { pending: true, ..GuardWalk::default() };
824 walker.walk(suffix, &mut suffix_walk);
825 has_tail_guard |= !suffix_walk.failed && !suffix_walk.pending;
826 }
827 }
828 if has_tail_guard {
829 state.cover(GuardCoverage::Callback, true);
830 }
831 GuardWalk {
832 coverage: state.coverage,
833 future_coverage: state.future_coverage,
834 ..GuardWalk::default()
835 }
836 }
837}
838
839struct Finder<'gcx, S, E> {
841 gcx: Gcx<'gcx>,
842 stmt_matches: S,
843 expr_matches: E,
844}
845
846impl<'gcx, S, E> Visit<'gcx> for Finder<'gcx, S, E>
847where
848 S: FnMut(&'gcx Stmt<'gcx>) -> bool,
849 E: FnMut(&'gcx Expr<'gcx>) -> bool,
850{
851 type BreakValue = ();
852
853 fn hir(&self) -> &'gcx Hir<'gcx> {
854 &self.gcx.hir
855 }
856
857 fn visit_stmt(&mut self, stmt: &'gcx Stmt<'gcx>) -> ControlFlow<()> {
858 if (self.stmt_matches)(stmt) { ControlFlow::Break(()) } else { self.walk_stmt(stmt) }
859 }
860
861 fn visit_expr(&mut self, expr: &'gcx Expr<'gcx>) -> ControlFlow<()> {
862 if (self.expr_matches)(expr) { ControlFlow::Break(()) } else { self.walk_expr(expr) }
863 }
864}
865
866#[derive(Clone, Copy, Default, PartialEq, Eq)]
871enum GuardCoverage {
872 #[default]
873 None,
874 Callback,
875 CodeLess,
876 CallbackOrCodeLess,
877}
878
879impl GuardCoverage {
880 fn is_covered(self) -> bool {
881 self != Self::None
882 }
883
884 const fn relies_on_code_length(self) -> bool {
885 matches!(self, Self::CodeLess | Self::CallbackOrCodeLess)
886 }
887
888 const fn merge_paths(self, other: Self) -> Self {
891 match (self, other) {
892 (Self::None, _) | (_, Self::None) => Self::None,
893 (Self::Callback, Self::Callback) => Self::Callback,
894 (Self::CodeLess, Self::CodeLess) => Self::CodeLess,
895 _ => Self::CallbackOrCodeLess,
896 }
897 }
898
899 const fn combine_guards(self, other: Self) -> Self {
902 match (self, other) {
903 (Self::Callback, _) | (_, Self::Callback) => Self::Callback,
904 (Self::CallbackOrCodeLess, _) | (_, Self::CallbackOrCodeLess) => {
905 Self::CallbackOrCodeLess
906 }
907 (Self::CodeLess, _) | (_, Self::CodeLess) => Self::CodeLess,
908 _ => Self::None,
909 }
910 }
911}
912
913#[derive(Clone, Default)]
918struct GuardWalk {
919 coverage: GuardCoverage,
921 future_coverage: GuardCoverage,
925 pending: bool,
926 failed: bool,
927 escaped: bool,
928}
929
930impl GuardWalk {
931 const fn cover(&mut self, coverage: GuardCoverage, future: bool) {
933 self.coverage = self.coverage.combine_guards(coverage);
934 if future {
935 self.future_coverage = self.future_coverage.combine_guards(coverage);
936 }
937 self.pending = false;
938 }
939
940 const fn retire(&mut self) {
944 self.failed |= self.pending;
945 self.coverage = GuardCoverage::None;
946 self.future_coverage = GuardCoverage::None;
947 }
948
949 fn escape(&mut self) {
951 self.failed |= self.pending;
952 self.escaped |= !self.coverage.is_covered();
953 }
954
955 fn retire_code_snapshots_if(&mut self, may_change_code: impl FnOnce() -> bool) {
958 let (coverage, future) =
959 (self.coverage.relies_on_code_length(), self.future_coverage.relies_on_code_length());
960 if (coverage || future) && may_change_code() {
961 if coverage {
962 self.coverage = GuardCoverage::None;
963 }
964 if future {
965 self.future_coverage = GuardCoverage::None;
966 }
967 }
968 }
969
970 const fn merge(self, other: Self) -> Self {
973 Self {
974 coverage: self.coverage.merge_paths(other.coverage),
975 future_coverage: self.future_coverage.merge_paths(other.future_coverage),
976 pending: self.pending || other.pending,
977 failed: self.failed || other.failed,
978 escaped: self.escaped || other.escaped,
979 }
980 }
981}
982
983struct GuardWalker<'a, 'gcx> {
1007 cx: Cx<'gcx>,
1008 recipient: VariableId,
1009 token: VariableId,
1010 delegations: &'a [FunctionId],
1012 unstable_code_delegations: &'a [FunctionId],
1013 seen: &'a mut Vec<FunctionId>,
1014}
1015
1016impl<'gcx> GuardWalker<'_, 'gcx> {
1017 fn walk(&mut self, stmts: &'gcx [Stmt<'gcx>], walk: &mut GuardWalk) {
1018 let cx = self.cx;
1019 for stmt in stmts {
1022 let guard = match &stmt.kind {
1023 StmtKind::Expr(expr) => self.guard_expr_coverage(expr),
1024 _ => GuardCoverage::None,
1025 };
1026 match &stmt.kind {
1027 StmtKind::Block(block) | StmtKind::UncheckedBlock(block) => {
1028 self.walk(block.stmts, walk);
1029 }
1030 StmtKind::Expr(expr) if guard.is_covered() => {
1031 if self.mutates(stmt) {
1037 walk.retire();
1038 } else if cx.may_return(stmt) {
1039 walk.escape();
1040 } else if guard.relies_on_code_length()
1041 && self.guard_extra_args_may_change_account_code(expr)
1042 {
1043 if walk.future_coverage.relies_on_code_length() {
1044 walk.future_coverage = GuardCoverage::None;
1045 }
1046 } else {
1047 walk.cover(guard, guard == GuardCoverage::CodeLess);
1048 }
1049 }
1050 StmtKind::If(cond, then, otherwise) => {
1051 let condition_mutates = cx.expr_mutates_var(cond, self.recipient)
1055 || cx.expr_mutates_var(cond, self.token);
1056 if condition_mutates {
1057 walk.retire();
1058 }
1059 if walk.future_coverage.relies_on_code_length()
1060 && self.may_change_account_code(slice::from_ref(stmt), Some(cond))
1061 {
1062 walk.future_coverage = GuardCoverage::None;
1063 }
1064 if cx.expr_contains_frame_ending_assembly(cond) {
1065 walk.escape();
1066 }
1067 let refusal_then = !condition_mutates
1071 && self.is_hook_comparison(cond, BinOpKind::Ne)
1072 && cx.branch_always_reverts(then);
1073 let refusal_else = !condition_mutates
1074 && self.is_hook_comparison(cond, BinOpKind::Eq)
1075 && otherwise.is_some_and(|otherwise| cx.branch_always_reverts(otherwise));
1076 if refusal_then || refusal_else {
1077 walk.cover(GuardCoverage::Callback, false);
1078 let accepted = if refusal_then { *otherwise } else { Some(*then) };
1079 if let Some(accepted) = accepted {
1080 self.walk(slice::from_ref(accepted), walk);
1081 }
1082 continue;
1083 }
1084 let mut then_walk = walk.clone();
1089 let mut else_walk = walk.clone();
1090 if self.is_code_length_test(cond, true) {
1091 else_walk.cover(GuardCoverage::CodeLess, true);
1092 } else if self.is_code_length_test(cond, false) {
1093 then_walk.cover(GuardCoverage::CodeLess, true);
1094 }
1095 self.walk(slice::from_ref(then), &mut then_walk);
1096 if let Some(otherwise) = otherwise {
1097 self.walk(slice::from_ref(otherwise), &mut else_walk);
1098 }
1099 *walk = then_walk.merge(else_walk);
1100 }
1101 _ => {
1102 if self.mutates(stmt) {
1105 walk.retire();
1106 }
1107 if walk.future_coverage.relies_on_code_length()
1112 && self.may_change_account_code(slice::from_ref(stmt), None)
1113 {
1114 walk.future_coverage = GuardCoverage::None;
1115 }
1116 let delegations = self.delegations;
1120 if !walk.future_coverage.is_covered()
1121 && cx.any_in_stmts(
1122 slice::from_ref(stmt),
1123 |_| false,
1124 |expr| {
1125 cx.resolved_callee(expr).is_some_and(|id| delegations.contains(&id))
1126 },
1127 )
1128 {
1129 walk.pending = true;
1130 }
1131 if cx.may_return(stmt) {
1132 walk.escape();
1133 }
1134 }
1135 }
1136 }
1137 }
1138
1139 fn mutates(&self, stmt: &'gcx Stmt<'gcx>) -> bool {
1140 self.cx.mutates_var(stmt, self.recipient) || self.cx.mutates_var(stmt, self.token)
1141 }
1142
1143 fn may_change_account_code(
1146 &self,
1147 stmts: &'gcx [Stmt<'gcx>],
1148 expr: Option<&'gcx Expr<'gcx>>,
1149 ) -> bool {
1150 let (delegations, unstable) = (self.delegations, self.unstable_code_delegations);
1151 let mut seen = Vec::new();
1152 match expr {
1153 Some(expr) => {
1154 self.cx.expr_may_change_account_code(expr, delegations, unstable, &mut seen)
1155 }
1156 None => self.cx.stmts_may_change_account_code(stmts, delegations, unstable, &mut seen),
1157 }
1158 }
1159
1160 fn guard_expr_coverage(&mut self, expr: &'gcx Expr<'gcx>) -> GuardCoverage {
1166 let expr = expr.peel_parens();
1167 let ExprKind::Call(callee, args) = &expr.kind else { return GuardCoverage::None };
1168 if is_require_or_assert(self.cx.gcx, callee) {
1169 return args
1170 .exprs()
1171 .next()
1172 .map_or(GuardCoverage::None, |cond| self.acceptance_coverage(cond));
1173 }
1174 let Some(function_id) = self.cx.resolved_internal_callee(expr) else {
1175 return GuardCoverage::None;
1176 };
1177 let Some((recipient, token)) =
1178 self.cx.bound_guard_parameters(function_id, args, self.recipient, self.token)
1179 else {
1180 return GuardCoverage::None;
1181 };
1182 self.cx.body_guards(function_id, recipient, token, self.seen)
1183 }
1184
1185 fn guard_extra_args_may_change_account_code(&self, expr: &'gcx Expr<'gcx>) -> bool {
1190 let ExprKind::Call(callee, args) = &expr.peel_parens().kind else { return false };
1191 is_require_or_assert(self.cx.gcx, callee)
1192 && args.exprs().skip(1).any(|arg| self.may_change_account_code(&[], Some(arg)))
1193 }
1194
1195 fn acceptance_coverage(&self, cond: &'gcx Expr<'gcx>) -> GuardCoverage {
1199 let cond = cond.peel_parens();
1200 if self.is_code_length_test(cond, false) {
1201 return GuardCoverage::CodeLess;
1202 }
1203 if self.is_hook_comparison(cond, BinOpKind::Eq) {
1204 return GuardCoverage::Callback;
1205 }
1206 let ExprKind::Binary(lhs, op, rhs) = &cond.kind else { return GuardCoverage::None };
1207 let accepts = |skip, check| {
1208 self.is_code_length_test(skip, false) && self.is_hook_comparison(check, BinOpKind::Eq)
1209 };
1210 if op.kind == BinOpKind::Or && (accepts(lhs, rhs) || accepts(rhs, lhs)) {
1211 GuardCoverage::CallbackOrCodeLess
1212 } else {
1213 GuardCoverage::None
1214 }
1215 }
1216
1217 fn is_hook_call_on(&self, expr: &'gcx Expr<'gcx>) -> bool {
1221 let expr = expr.peel_parens();
1222 let Some((callee, args, _)) = expr.as_call() else { return false };
1223 let ExprKind::Member(receiver, _) = &callee.peel_parens().kind else { return false };
1224 let Some(function_id) = self.cx.resolved_callee(expr) else { return false };
1225 self.cx.is_receiver_hook(function_id)
1226 && underlying_var(self.cx.gcx, receiver) == Some(self.recipient)
1227 && self.cx.arg(function_id, args, 2).and_then(|expr| underlying_var(self.cx.gcx, expr))
1228 == Some(self.token)
1229 }
1230
1231 fn is_hook_comparison(&self, expr: &'gcx Expr<'gcx>, want: BinOpKind) -> bool {
1236 let ExprKind::Binary(lhs, op, rhs) = &expr.peel_parens().kind else { return false };
1237 let compares = |hook, answer| {
1238 self.is_hook_call_on(hook)
1239 && !self.is_hook_call_on(answer)
1240 && self.cx.is_received_selector(answer)
1241 };
1242 op.kind == want && (compares(lhs, rhs) || compares(rhs, lhs))
1243 }
1244
1245 fn is_code_length_test(&self, expr: &'gcx Expr<'gcx>, has_code: bool) -> bool {
1249 let ExprKind::Binary(lhs, op, rhs) = &expr.peel_parens().kind else { return false };
1250 let is_code_length = |expr: &Expr<'_>| {
1251 let ExprKind::Member(code, length) = &expr.peel_parens().kind else { return false };
1252 let ExprKind::Member(base, member) = &code.peel_parens().kind else { return false };
1253 length.as_str() == "length"
1254 && member.as_str() == "code"
1255 && underlying_var(self.cx.gcx, base) == Some(self.recipient)
1256 };
1257 let literal = |expr: &Expr<'_>| match &expr.peel_parens().kind {
1258 ExprKind::Lit(lit) => match &lit.kind {
1259 LitKind::Number(value) => u8::try_from(*value).ok(),
1260 _ => None,
1261 },
1262 _ => None,
1263 };
1264 let (bound, flipped) = if is_code_length(lhs) {
1265 (literal(rhs), false)
1266 } else if is_code_length(rhs) {
1267 (literal(lhs), true)
1268 } else {
1269 return false;
1270 };
1271 let Some(bound) = bound else { return false };
1272 match (has_code, op.kind, flipped) {
1276 (true, BinOpKind::Ne, _)
1277 | (true, BinOpKind::Gt, false)
1278 | (true, BinOpKind::Lt, true)
1279 | (false, BinOpKind::Eq, _)
1280 | (false, BinOpKind::Le, false)
1281 | (false, BinOpKind::Ge, true) => bound == 0,
1282 (true, BinOpKind::Ge, false)
1283 | (true, BinOpKind::Le, true)
1284 | (false, BinOpKind::Lt, false)
1285 | (false, BinOpKind::Gt, true) => bound == 1,
1286 _ => false,
1287 }
1288 }
1289}
1290
1291#[derive(Clone, Copy)]
1293enum SelectorEncoding {
1294 Literal,
1295 Integer(u16),
1296 FixedBytes(u8),
1297}
1298
1299const ERC721_RECEIVED: u64 = 0x150b_7a02;
1301
1302fn is_canonical_erc721(name: &str) -> bool {
1308 matches!(
1309 name,
1310 "ERC721" | "ERC721Upgradeable" | "ERC721Consecutive" | "ERC721ConsecutiveUpgradeable"
1311 )
1312}
1313
1314fn named(function: &hir::Function<'_>, name: &str) -> bool {
1315 function.name.is_some_and(|n| n.as_str() == name)
1316}
1317
1318const fn is_internal(function: &hir::Function<'_>) -> bool {
1319 matches!(function.visibility, Visibility::Internal | Visibility::Private)
1320}
1321
1322const fn is_assembly(stmt: &Stmt<'_>) -> bool {
1323 matches!(stmt.kind, StmtKind::AssemblyBlock(_))
1324}
1325
1326fn keeps_its_value(gcx: Gcx<'_>, variable: VariableId) -> bool {
1330 let variable = gcx.hir.variable(variable);
1331 !variable.kind.is_state() || variable.mutability.is_some()
1332}
1333
1334fn assigns_to(gcx: Gcx<'_>, expr: &Expr<'_>, var: VariableId) -> bool {
1337 let Some(target) = write_target(expr) else { return false };
1338 let mut hit = false;
1339 for_each_lhs_var(gcx, target, &mut |vid| hit |= vid == var);
1340 hit
1341}
1342
1343fn is_revert_call(gcx: Gcx<'_>, expr: &Expr<'_>) -> bool {
1345 let ExprKind::Call(callee, args) = &expr.peel_parens().kind else { return false };
1346 is_builtin(gcx, callee, kw::Revert)
1347 || (is_require_or_assert(gcx, callee) && args.exprs().next().is_some_and(is_literal_false))
1348}
1349
1350fn modifier_body_sides<'gcx>(
1353 stmts: &'gcx [Stmt<'gcx>],
1354) -> Option<(&'gcx [Stmt<'gcx>], &'gcx [Stmt<'gcx>])> {
1355 let placeholders =
1356 stmts.iter().enumerate().filter(|(_, stmt)| matches!(stmt.kind, StmtKind::Placeholder));
1357 let index = unique(placeholders.map(|(index, _)| index))?;
1358 Some((&stmts[..index], &stmts[index + 1..]))
1359}