extern crate rustc_hir;
extern crate rustc_span;
use std::collections::HashMap;
use rustc_hir::Block;
use rustc_hir::Body;
use rustc_hir::Expr;
use rustc_hir::ExprKind;
use rustc_hir::HirId;
use rustc_hir::MatchSource;
use rustc_hir::PatKind;
use rustc_hir::QPath;
use rustc_hir::Stmt;
use rustc_hir::StmtKind;
use rustc_hir::def::Res;
use rustc_hir::intravisit::FnKind;
use rustc_lint::LateContext;
use rustc_lint::LateLintPass;
use rustc_span::Span;
use crate::diagnostics;
use crate::shared;
crate::declare_late_lint! {
pub REQUIRE_GUARDED_FULL_BALANCE_DRAIN,
Warn,
"full-balance drains should be gated by a pause or circuit-breaker guard"
}
const DRAIN_METHODS: &[&str] = &["send", "send_owned"];
const CLOSE_METHODS: &[&str] = &[
"zeroed",
"close",
"close_with_recipient",
"close_account_zeroed",
];
const GUARD_TERMS: &[&str] = &[
"pause", "cap", "circuit", "halt", "guard", "limit", "throttle",
];
const TARGET_NEEDLES: &[&str] = &["process", "process_instruction", "instruction"];
#[derive(Debug, Clone, Copy)]
struct DrainFacts {
full_balance: bool,
guarded: bool,
closing: bool,
}
#[derive(Debug, Clone)]
struct SeenGuard {
scopes: Vec<u32>,
receiver: Option<String>,
is_close: bool,
}
#[derive(Default)]
struct DrainAnalyzer<'tcx> {
scopes: Vec<u32>,
next_scope: u32,
guards: Vec<SeenGuard>,
initializers: HashMap<HirId, &'tcx Expr<'tcx>>,
drains: HashMap<Span, DrainFacts>,
order: Vec<Span>,
}
impl<'tcx> DrainAnalyzer<'tcx> {
fn analyze(body: &'tcx Body<'tcx>) -> Self {
let mut analyzer = Self::default();
analyzer.visit_expr(body.value);
analyzer
}
fn has_dominant(&self, keep: impl Fn(&SeenGuard) -> bool) -> bool {
self.guards
.iter()
.any(|guard| keep(guard) && self.scopes.starts_with(&guard.scopes))
}
fn resolves_to_full_balance(&self, expr: &Expr<'_>, receiver: &str) -> bool {
match &expr.kind {
ExprKind::MethodCall(segment, inner_receiver, arguments, _) => {
segment.ident.name.as_str() == "lamports"
&& arguments.is_empty()
&& shared::expression_identity(inner_receiver).as_deref() == Some(receiver)
}
ExprKind::Path(QPath::Resolved(_, path)) => {
match path.res {
Res::Local(binding) => {
self.initializers.get(&binding).is_some_and(|initializer| {
self.resolves_to_full_balance(initializer, receiver)
})
}
_ => false,
}
}
ExprKind::DropTemps(inner)
| ExprKind::Use(inner, _)
| ExprKind::AddrOf(_, _, inner)
| ExprKind::Unary(_, inner) => self.resolves_to_full_balance(inner, receiver),
_ => false,
}
}
fn record_guard(&mut self, method: &str, receiver: Option<String>) {
let lowercase = method.to_ascii_lowercase();
let is_guard = GUARD_TERMS.iter().any(|term| lowercase.contains(term));
let is_close = CLOSE_METHODS.contains(&method);
if is_guard || is_close {
self.guards.push(SeenGuard {
scopes: self.scopes.clone(),
receiver,
is_close,
});
}
}
fn visit_branch(&mut self, expr: &'tcx Expr<'tcx>) {
let id = self.next_scope;
self.next_scope += 1;
self.scopes.push(id);
self.visit_expr(expr);
self.scopes.pop();
}
fn visit_branch_block(&mut self, block: &'tcx Block<'tcx>) {
let id = self.next_scope;
self.next_scope += 1;
self.scopes.push(id);
self.visit_block(block);
self.scopes.pop();
}
fn visit_block(&mut self, block: &'tcx Block<'tcx>) {
for statement in block.stmts {
self.visit_stmt(statement);
}
if let Some(tail) = block.expr {
self.visit_expr(tail);
}
}
fn visit_stmt(&mut self, statement: &'tcx Stmt<'tcx>) {
if let StmtKind::Let(local) = &statement.kind {
if let (PatKind::Binding(_, binding, ..), Some(initializer)) =
(&local.pat.kind, local.init)
{
self.initializers.insert(*binding, initializer);
}
if let Some(initializer) = local.init {
self.visit_expr(initializer);
}
if let Some(else_block) = local.els {
self.visit_branch_block(else_block);
}
} else if let StmtKind::Expr(expr) | StmtKind::Semi(expr) = &statement.kind {
self.visit_expr(expr);
}
}
fn visit_expr(&mut self, expr: &'tcx Expr<'tcx>) {
match &expr.kind {
ExprKind::MethodCall(segment, receiver, arguments, _) => {
let method = segment.ident.name.as_str();
let receiver_identity = shared::expression_identity(receiver);
if DRAIN_METHODS.contains(&method) && arguments.len() == 3 {
let full_balance = receiver_identity.as_deref().is_some_and(|receiver| {
arguments
.get(1)
.is_some_and(|amount| self.resolves_to_full_balance(amount, receiver))
});
let guarded = self.has_dominant(|guard| !guard.is_close);
let closing = self.has_dominant(|guard| {
guard.is_close
&& guard.receiver.is_some()
&& guard.receiver == receiver_identity
});
self.drains.insert(
expr.span,
DrainFacts {
full_balance,
guarded,
closing,
},
);
self.order.push(expr.span);
} else {
self.record_guard(method, receiver_identity);
}
self.visit_expr(receiver);
for argument in *arguments {
self.visit_expr(argument);
}
}
ExprKind::Block(block, _) => self.visit_block(block),
ExprKind::If(condition, then, otherwise) => {
self.visit_expr(condition);
self.visit_branch(then);
if let Some(otherwise) = otherwise {
self.visit_branch(otherwise);
}
}
ExprKind::Match(scrutinee, arms, source) => {
let try_desugar = matches!(source, MatchSource::TryDesugar(_));
self.visit_expr(scrutinee);
for arm in *arms {
if let Some(guard) = arm.guard {
self.visit_branch(guard);
}
if try_desugar {
self.visit_expr(arm.body);
} else {
self.visit_branch(arm.body);
}
}
}
ExprKind::Loop(block, ..) => self.visit_branch_block(block),
ExprKind::Closure(closure) => {
let _ = closure;
}
ExprKind::Call(callee, arguments) => {
self.visit_expr(callee);
for argument in *arguments {
self.visit_expr(argument);
}
}
ExprKind::Binary(_, left, right) => {
self.visit_expr(left);
self.visit_expr(right);
}
ExprKind::Assign(left, right, _) | ExprKind::AssignOp(_, left, right) => {
self.visit_expr(left);
self.visit_expr(right);
}
ExprKind::Index(base, index, _) => {
self.visit_expr(base);
self.visit_expr(index);
}
ExprKind::Let(let_expr) => self.visit_expr(let_expr.init),
ExprKind::Tup(expressions) | ExprKind::Array(expressions) => {
for expression in *expressions {
self.visit_expr(expression);
}
}
ExprKind::Struct(_, fields, tail) => {
for field in *fields {
self.visit_expr(field.expr);
}
if let rustc_hir::StructTailExpr::Base(base) = tail {
self.visit_expr(base);
}
}
ExprKind::Ret(Some(inner)) | ExprKind::Break(_, Some(inner)) => self.visit_expr(inner),
ExprKind::Unary(_, inner)
| ExprKind::Use(inner, _)
| ExprKind::Cast(inner, _)
| ExprKind::Type(inner, _)
| ExprKind::DropTemps(inner)
| ExprKind::AddrOf(_, _, inner)
| ExprKind::Field(inner, _)
| ExprKind::Repeat(inner, _)
| ExprKind::Yield(inner, _)
| ExprKind::Become(inner)
| ExprKind::UnsafeBinderCast(_, inner, _) => self.visit_expr(inner),
_ => {}
}
}
}
impl<'tcx> LateLintPass<'tcx> for RequireGuardedFullBalanceDrain {
fn check_fn(
&mut self,
cx: &LateContext<'tcx>,
_: FnKind<'tcx>,
_: &'tcx rustc_hir::FnDecl<'tcx>,
body: &'tcx Body<'tcx>,
_: Span,
def_id: rustc_hir::def_id::LocalDefId,
) {
let def_path = cx.tcx.def_path_str(def_id.to_def_id());
if shared::should_skip_def_path(&def_path)
|| !shared::def_path_matches(&def_path, TARGET_NEEDLES)
{
return;
}
let analyzer = DrainAnalyzer::analyze(body);
for span in &analyzer.order {
let Some(facts) = analyzer.drains.get(span) else {
continue;
};
if !facts.full_balance || facts.guarded || facts.closing {
continue;
}
diagnostics::emit(cx, REQUIRE_GUARDED_FULL_BALANCE_DRAIN, |diag| {
diag.span(*span);
diag.primary_message(
"an instruction path can sweep an account's entire balance in one call",
);
diag.help(
"gate full-balance sweeps behind a pause or circuit-breaker check (a pause \
flag plus a per-window withdrawal cap bounds a compromised key's blast \
radius)",
);
diag.help(
"if this drain is an account-close path, use `close_account_zeroed` so \
`require_zeroed_before_close` covers the stale-data risk too",
);
diag.help(shared::CONTROL_FLOW_LIMITATION_HELP);
});
}
}
}