1use super::{
5 branch_always_exits, is_require_or_assert, is_sender_member, lhs_local_var, loop_stmts,
6 stmt_expr, tuple_elems, underlying_var, visit_stmts,
7};
8use solar::sema::{
9 Gcx,
10 hir::{self, BinOpKind, Expr, ExprKind, FunctionId, Stmt, StmtKind, UnOpKind, VariableId},
11};
12use std::{collections::HashSet, iter, ops::ControlFlow};
13
14pub fn is_protected<'gcx>(gcx: Gcx<'gcx>, func_id: FunctionId) -> bool {
16 modifiers_and_self(gcx, func_id).any(|id| has_access_guard(gcx, id, &mut HashSet::new()))
17}
18
19pub fn modifiers_and_self<'gcx>(
21 gcx: Gcx<'gcx>,
22 func_id: FunctionId,
23) -> impl Iterator<Item = FunctionId> + 'gcx {
24 gcx.hir
25 .function(func_id)
26 .modifiers
27 .iter()
28 .filter_map(move |modifier| {
29 gcx.hir.function(func_id).contract.map_or_else(
30 || modifier.id.as_function(),
31 |contract| gcx.resolve_modifier_target(contract, modifier),
32 )
33 })
34 .chain(iter::once(func_id))
35}
36
37pub fn has_access_guard<'gcx>(
41 gcx: Gcx<'gcx>,
42 func_id: FunctionId,
43 seen: &mut HashSet<FunctionId>,
44) -> bool {
45 if !seen.insert(func_id) {
46 return false;
47 }
48 let func = gcx.hir.function(func_id);
49 match func.body {
50 Some(body) => for_each_guard(gcx, body, seen, &mut |_| ControlFlow::Break(())).is_break(),
51 None => looks_like_access_control(func),
52 }
53}
54
55pub fn guard_vars<'gcx>(gcx: Gcx<'gcx>, func_id: FunctionId) -> HashSet<VariableId> {
57 let mut out = HashSet::new();
58 for id in modifiers_and_self(gcx, func_id) {
59 let Some(body) = gcx.hir.function(id).body else { continue };
60 let mut seen = HashSet::from([id]);
61 let _ = for_each_guard(gcx, body, &mut HashSet::from([id]), &mut |guard| {
62 match guard {
63 Guard::Check(cond) => expr_state_vars(gcx, cond, &mut seen, &mut out),
64 Guard::Call(callee_id) => function_state_vars(gcx, callee_id, &mut seen, &mut out),
65 }
66 ControlFlow::Continue(())
67 });
68 }
69 out
70}
71
72pub fn looks_like_access_control(func: &hir::Function<'_>) -> bool {
75 let Some(name) = func.name else { return false };
76 if !func.returns.is_empty() {
77 return false;
78 }
79 let lower = name.as_str().to_ascii_lowercase();
80 matches!(lower.as_str(), "auth" | "requiresauth" | "restricted")
81 || ["only", "check", "_check"].iter().any(|prefix| {
82 ["admin", "guardian", "manager", "owner", "role"]
83 .iter()
84 .any(|role| lower.starts_with(&format!("{prefix}{role}")))
85 })
86}
87
88pub fn access_check_polarity<'gcx>(
93 gcx: Gcx<'gcx>,
94 expr: &Expr<'_>,
95 aliases: &HashSet<VariableId>,
96) -> Option<bool> {
97 let is_check = |sender: &Expr<'_>, authority: &Expr<'_>| {
98 expr_reads_sender(gcx, sender, &mut HashSet::new(), aliases)
99 && expr_reads_state(gcx, authority)
100 };
101 match &expr.peel_parens().kind {
102 ExprKind::Unary(op, inner) if op.kind == UnOpKind::Not => {
103 access_check_polarity(gcx, inner, aliases).map(|polarity| !polarity)
104 }
105 ExprKind::Binary(lhs, op, rhs) if matches!(op.kind, BinOpKind::And | BinOpKind::Or) => {
106 let dominant = op.kind == BinOpKind::And;
109 let lhs = access_check_polarity(gcx, lhs, aliases);
110 let rhs = access_check_polarity(gcx, rhs, aliases);
111 if lhs == Some(dominant) || rhs == Some(dominant) {
112 Some(dominant)
113 } else if lhs == Some(!dominant) && rhs == Some(!dominant) {
114 Some(!dominant)
115 } else {
116 None
117 }
118 }
119 ExprKind::Binary(lhs, op, rhs)
120 if matches!(op.kind, BinOpKind::Eq | BinOpKind::Ne)
121 && (is_check(lhs, rhs) || is_check(rhs, lhs)) =>
122 {
123 Some(op.kind == BinOpKind::Eq)
124 }
125 _ => is_check(expr, expr).then_some(true),
126 }
127}
128
129fn update_sender_aliases<'gcx>(
133 gcx: Gcx<'gcx>,
134 stmt: &Stmt<'gcx>,
135 aliases: &mut HashSet<VariableId>,
136) {
137 let reads_sender = |value: Option<&Expr<'_>>, aliases: &HashSet<VariableId>| {
138 value.is_some_and(|value| expr_reads_sender(gcx, value, &mut HashSet::new(), aliases))
139 };
140 let updates: Vec<(VariableId, bool)> = match stmt.kind {
143 StmtKind::DeclSingle(var_id) => match gcx.hir.variable(var_id).initializer {
144 Some(value) => vec![(var_id, reads_sender(Some(value), aliases))],
145 None => return,
146 },
147 StmtKind::DeclMulti(var_ids, value) => var_ids
148 .iter()
149 .enumerate()
150 .filter_map(|(i, var_id)| {
151 let value =
152 tuple_elems(value).map_or(Some(value), |elems| elems.get(i).copied().flatten());
153 var_id.map(|var_id| (var_id, reads_sender(value, aliases)))
154 })
155 .collect(),
156 StmtKind::Expr(expr) => match &expr.peel_parens().kind {
157 ExprKind::Assign(lhs, _, rhs) => {
158 let mut updates = Vec::new();
159 collect_sender_alias_updates(gcx, lhs, Some(rhs), aliases, &mut updates);
160 updates
161 }
162 _ => return,
163 },
164 _ => return,
165 };
166 for (var_id, reads_sender) in updates.into_iter().rev() {
169 if reads_sender {
170 aliases.insert(var_id);
171 } else {
172 aliases.remove(&var_id);
173 }
174 }
175}
176
177fn collect_sender_alias_updates(
180 gcx: Gcx<'_>,
181 lhs: &Expr<'_>,
182 rhs: Option<&Expr<'_>>,
183 aliases: &HashSet<VariableId>,
184 updates: &mut Vec<(VariableId, bool)>,
185) {
186 if let Some(lhs_elems) = tuple_elems(lhs) {
187 for (i, lhs) in lhs_elems.iter().enumerate() {
188 let Some(lhs) = lhs else { continue };
189 let rhs = rhs.and_then(|rhs| {
190 tuple_elems(rhs).map_or(Some(rhs), |elems| elems.get(i).copied().flatten())
191 });
192 collect_sender_alias_updates(gcx, lhs, rhs, aliases, updates);
193 }
194 } else if let Some(var_id) = lhs_local_var(gcx, lhs) {
195 let reads_sender =
196 rhs.is_some_and(|rhs| expr_reads_sender(gcx, rhs, &mut HashSet::new(), aliases));
197 updates.push((var_id, reads_sender));
198 }
199}
200
201pub fn expr_reads_sender<'gcx>(
204 gcx: Gcx<'gcx>,
205 expr: &Expr<'_>,
206 seen: &mut HashSet<FunctionId>,
207 aliases: &HashSet<VariableId>,
208) -> bool {
209 expr.visit(&mut |e| {
210 let reads = is_sender_member(gcx, e)
211 || underlying_var(gcx, e).is_some_and(|v| aliases.contains(&v))
212 || matches!(&e.kind, ExprKind::Call(callee, ..)
213 if matches!(callee.peel_parens().kind, ExprKind::Ident(_))
214 && gcx.resolved_function(callee).is_some_and(|id| function_reads_sender(gcx, id, seen)));
215 if reads { ControlFlow::Break(()) } else { ControlFlow::Continue(()) }
216 })
217 .is_break()
218}
219
220pub fn function_reads_sender<'gcx>(
222 gcx: Gcx<'gcx>,
223 func_id: FunctionId,
224 seen: &mut HashSet<FunctionId>,
225) -> bool {
226 seen.insert(func_id)
227 && gcx.hir.function(func_id).body.is_some_and(|body| {
228 visit_stmts(&gcx.hir, body.stmts, |stmt| {
229 let reads = stmt_expr(&gcx.hir, stmt)
230 .is_some_and(|expr| expr_reads_sender(gcx, expr, seen, &HashSet::new()));
231 if reads { ControlFlow::Break(()) } else { ControlFlow::Continue(()) }
232 })
233 .is_break()
234 })
235}
236
237pub fn expr_state_vars<'gcx>(
239 gcx: Gcx<'gcx>,
240 expr: &Expr<'_>,
241 seen: &mut HashSet<FunctionId>,
242 out: &mut HashSet<VariableId>,
243) {
244 let _ = expr.visit(&mut |e| {
245 if let Some(var_id) = underlying_var(gcx, e)
246 && gcx.hir.variable(var_id).kind.is_state()
247 {
248 out.insert(var_id);
249 }
250 if let ExprKind::Call(callee, ..) = &e.kind
251 && matches!(callee.peel_parens().kind, ExprKind::Ident(_))
252 && let Some(callee_id) = gcx.resolved_function(callee)
253 {
254 function_state_vars(gcx, callee_id, seen, out);
255 }
256 ControlFlow::<()>::Continue(())
257 });
258}
259
260pub fn function_state_vars<'gcx>(
262 gcx: Gcx<'gcx>,
263 func_id: FunctionId,
264 seen: &mut HashSet<FunctionId>,
265 out: &mut HashSet<VariableId>,
266) {
267 if seen.insert(func_id)
268 && let Some(body) = gcx.hir.function(func_id).body
269 {
270 let _ = visit_stmts(&gcx.hir, body.stmts, |stmt| {
271 if let Some(expr) = stmt_expr(&gcx.hir, stmt) {
272 expr_state_vars(gcx, expr, seen, out);
273 }
274 ControlFlow::Continue(())
275 });
276 }
277}
278
279fn expr_reads_state<'gcx>(gcx: Gcx<'gcx>, expr: &Expr<'_>) -> bool {
280 let mut vars = HashSet::new();
281 expr_state_vars(gcx, expr, &mut HashSet::new(), &mut vars);
282 !vars.is_empty()
283}
284
285enum Guard<'a> {
287 Check(&'a Expr<'a>),
289 Call(FunctionId),
291}
292
293fn for_each_guard<'gcx>(
296 gcx: Gcx<'gcx>,
297 body: hir::Block<'gcx>,
298 seen: &mut HashSet<FunctionId>,
299 f: &mut impl FnMut(Guard<'_>) -> ControlFlow<()>,
300) -> ControlFlow<()> {
301 let mut stmts = Vec::new();
302 let _ = dominating_stmts(body.stmts, &mut stmts);
303 let mut aliases = HashSet::new();
307 for stmt in stmts {
308 if let StmtKind::If(cond, then_stmt, else_stmt) = stmt.kind {
309 let exits = match access_check_polarity(gcx, cond, &aliases) {
310 Some(false) => branch_always_exits(gcx, then_stmt),
311 Some(true) => else_stmt.is_some_and(|expr| branch_always_exits(gcx, expr)),
312 None => false,
313 };
314 if exits {
315 f(Guard::Check(cond))?;
316 }
317 continue;
318 }
319 update_sender_aliases(gcx, stmt, &mut aliases);
320 let Some(expr) = stmt_expr(&gcx.hir, stmt) else { continue };
321 expr.visit(&mut |e| {
322 match &e.kind {
323 ExprKind::Call(callee, args) if is_require_or_assert(gcx, callee) => {
324 if let Some(cond) = args.exprs().next()
325 && access_check_polarity(gcx, cond, &aliases) == Some(true)
326 {
327 f(Guard::Check(cond))?;
328 }
329 }
330 ExprKind::Call(callee, ..)
331 if matches!(callee.peel_parens().kind, ExprKind::Ident(_)) =>
332 {
333 if let Some(callee_id) = gcx.resolved_function(callee)
334 && (looks_like_access_control(gcx.hir.function(callee_id))
335 || has_access_guard(gcx, callee_id, seen))
336 {
337 f(Guard::Call(callee_id))?;
338 }
339 }
340 _ => {}
341 }
342 ControlFlow::Continue(())
343 })?;
344 }
345 ControlFlow::Continue(())
346}
347
348fn dominating_stmts<'gcx>(
351 stmts: impl IntoIterator<Item = &'gcx Stmt<'gcx>>,
352 out: &mut Vec<&'gcx Stmt<'gcx>>,
353) -> ControlFlow<()> {
354 for stmt in stmts {
355 match stmt.kind {
356 StmtKind::Placeholder => return ControlFlow::Break(()),
357 StmtKind::Block(block) | StmtKind::UncheckedBlock(block) => {
358 dominating_stmts(block.stmts, out)?;
359 }
360 StmtKind::Loop(block, source) => dominating_stmts(loop_stmts(block, source), out)?,
361 _ => out.push(stmt),
362 }
363 }
364 ControlFlow::Continue(())
365}