extern crate rustc_hir;
extern crate rustc_span;
use rustc_hir::intravisit::FnKind;
use rustc_lint::LateContext;
use rustc_lint::LateLintPass;
use crate::diagnostics;
use crate::shared;
crate::declare_late_lint! {
pub REQUIRE_CANONICAL_BUMP_BEFORE_PDA_WRITE,
Deny,
"explicit PDA bumps must be proven canonical before use"
}
impl<'tcx> LateLintPass<'tcx> for RequireCanonicalBumpBeforePdaWrite {
fn check_fn(
&mut self,
cx: &LateContext<'tcx>,
_: FnKind<'tcx>,
_: &'tcx rustc_hir::FnDecl<'tcx>,
body: &'tcx rustc_hir::Body<'tcx>,
_: rustc_span::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, &["process", "instruction"])
{
return;
}
let mut facts = shared::collect_function_facts(cx, body);
let mut reassigned = std::collections::HashSet::new();
record_tuple_pattern_aliases(body, &mut facts.aliases, &mut reassigned);
for (index, call) in facts.calls.iter().enumerate() {
if call.method == "assert_stored_bump" {
if stored_bump_provenance_ok(&facts, call, &reassigned) {
continue;
}
diagnostics::emit(cx, REQUIRE_CANONICAL_BUMP_BEFORE_PDA_WRITE, |diag| {
diag.span(call.span);
diag.primary_message(
"assert_stored_bump was given a bump that was not parsed from the account",
);
diag.help(
"read the bump from the account's own state — for example capture it in \
the same `as_account`/`with_compact_account` parse that produced the \
other fields being validated — or use `assert_seeds()`, which performs \
that parse itself",
);
diag.help(shared::CONTROL_FLOW_LIMITATION_HELP);
});
continue;
}
if call.method != "assert_seeds_with_bump" {
continue;
}
let has_canonical_check = shared::has_prior_method_with_receiver_match(
&facts.calls,
index,
&["assert_canonical_bump", "assert_seeds"],
&call.receiver,
);
if has_canonical_check {
continue;
}
diagnostics::emit(cx, REQUIRE_CANONICAL_BUMP_BEFORE_PDA_WRITE, |diag| {
diag.span(call.span);
diag.primary_message(
"explicit PDA bump used without first proving the canonical address",
);
diag.help(
"call `account.assert_canonical_bump(seeds, program_id)?` before using an \
explicit bump in validation-only code, use `assert_seeds()`, or let \
`CreateProgramAccountWithBump` validate the canonical bump; \
`CreateProgramAccountWithUncheckedBump` skips that check on purpose",
);
diag.help(shared::CONTROL_FLOW_LIMITATION_HELP);
});
}
}
}
fn stored_bump_provenance_ok(
facts: &shared::FunctionFacts,
call: &shared::CallInfo,
reassigned: &std::collections::HashSet<rustc_hir::HirId>,
) -> bool {
let Some(account_identity) = call.args.first().and_then(Option::as_deref) else {
return false;
};
let mut identity = call.args.get(1).and_then(Option::as_deref);
let mut binding = call.arg_bindings.get(1).copied().flatten();
let mut visited = std::collections::HashSet::new();
loop {
if let Some(value) = identity {
if value == account_identity {
return true;
}
if let Some(rest) = value
.strip_prefix(account_identity)
.and_then(|suffix| suffix.strip_prefix('.'))
&& !rest.is_empty()
&& !rest.contains('.')
{
return true;
}
}
let Some(current) = binding else {
return false;
};
if !visited.insert(current) || reassigned.contains(¤t) {
return false;
};
let Some(alias) = facts.aliases.get(¤t) else {
return false;
};
identity = Some(alias.identity.as_str());
binding = alias.binding;
}
}
fn field_base_local_binding(expr: &rustc_hir::Expr<'_>) -> Option<rustc_hir::HirId> {
match expr.kind {
rustc_hir::ExprKind::Field(base, _) => shared::expression_local_binding(base),
_ => shared::expression_local_binding(expr),
}
}
fn record_tuple_pattern_aliases(
body: &rustc_hir::Body<'_>,
aliases: &mut std::collections::HashMap<rustc_hir::HirId, shared::AliasInfo>,
reassigned: &mut std::collections::HashSet<rustc_hir::HirId>,
) {
struct TupleAliases<'map> {
aliases: &'map mut std::collections::HashMap<rustc_hir::HirId, shared::AliasInfo>,
reassigned: &'map mut std::collections::HashSet<rustc_hir::HirId>,
}
impl<'hir> rustc_hir::intravisit::Visitor<'hir> for TupleAliases<'hir> {
fn visit_expr(&mut self, expr: &'hir rustc_hir::Expr<'hir>) {
if let rustc_hir::ExprKind::Assign(lhs, ..) = expr.kind
&& let Some(binding) =
field_base_local_binding(lhs).or_else(|| shared::expression_local_binding(lhs))
{
self.reassigned.insert(binding);
}
rustc_hir::intravisit::walk_expr(self, expr);
}
fn visit_stmt(&mut self, stmt: &'hir rustc_hir::Stmt<'hir>) {
if let rustc_hir::StmtKind::Let(local) = stmt.kind
&& let Some(init) = local.init
&& let rustc_hir::PatKind::Tuple(pat_elements, rest_position) = local.pat.kind
{
let mut tail = init;
while let rustc_hir::ExprKind::Block(block, _) = tail.kind {
match block.expr {
Some(next) => tail = next,
None => break,
}
}
let rest_position = rest_position.as_opt_usize();
if let rustc_hir::ExprKind::Tup(init_elements) = tail.kind
&& pat_elements.len() + usize::from(rest_position.is_some())
== init_elements.len()
{
for (pattern_index, pattern) in pat_elements.iter().enumerate() {
let element_index = match rest_position {
Some(rest) if pattern_index >= rest => pattern_index + 1,
_ => pattern_index,
};
let Some(element) = init_elements.get(element_index) else {
continue;
};
let rustc_hir::PatKind::Binding(_, binding, ..) = pattern.kind else {
continue;
};
if let Some(identity) = shared::expression_identity(element) {
self.aliases.insert(
binding,
shared::AliasInfo {
identity,
binding: field_base_local_binding(element),
},
);
}
}
}
}
rustc_hir::intravisit::walk_stmt(self, stmt);
}
}
rustc_hir::intravisit::walk_body(
&mut TupleAliases {
aliases,
reassigned,
},
body,
);
}