extern crate rustc_hir;
extern crate rustc_span;
use std::collections::HashSet;
use rustc_hir::intravisit::FnKind;
use rustc_lint::LateContext;
use rustc_lint::LateLintPass;
use rustc_lint::LintContext;
use crate::shared;
crate::declare_late_lint! {
pub REQUIRE_CONSISTENT_TOKEN_PROGRAM,
Deny,
"token operations in one instruction must share one program identity"
}
fn program_argument(call: &shared::CallInfo) -> Option<(usize, &str)> {
let index = match call.method.as_str() {
"as_token_mint_for_program" | "as_token_account_for_program" => 0,
"as_associated_token_account"
| "as_associated_token_account_checked"
| "assert_associated_token_address" => 2,
"invoke_with_program" => 0,
"invoke_signed_with_program" => 1,
_ => return None,
};
call.args
.get(index)
.and_then(Option::as_deref)
.map(|argument| (index, argument))
}
fn canonical_identity<'a>(
call: &'a shared::CallInfo,
argument_index: usize,
mut identity: &'a str,
facts: &'a shared::FunctionFacts,
) -> &'a str {
let mut visited = HashSet::new();
let mut binding = call.arg_bindings.get(argument_index).copied().flatten();
while let Some(current) = binding
&& visited.insert(current)
{
if facts
.assignments
.iter()
.any(|assignment| assignment.identity == identity)
{
break;
}
let Some(alias) = facts.aliases.get(¤t) else {
break;
};
identity = &alias.identity;
binding = alias.binding;
}
identity
}
impl<'tcx> LateLintPass<'tcx> for RequireConsistentTokenProgram {
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 facts = shared::collect_function_facts(cx, body);
let mut expected = None;
let mut previous_usage = None;
let mut reported_reassignments = HashSet::new();
for call in &facts.calls {
let Some((argument_index, argument)) = program_argument(call) else {
continue;
};
let identity = canonical_identity(call, argument_index, argument, &facts);
let Some(first) = expected else {
expected = Some(identity);
previous_usage = Some(call);
continue;
};
if identity == first {
let was_reassigned = previous_usage.is_some_and(|previous: &shared::CallInfo| {
facts.assignments.iter().any(|assignment| {
assignment.identity == identity
&& previous.span.hi() <= assignment.span.lo()
&& assignment.span.hi() <= call.span.lo()
})
});
if was_reassigned && reported_reassignments.insert(identity.to_string()) {
cx.lint(REQUIRE_CONSISTENT_TOKEN_PROGRAM, |diag| {
diag.span(call.span);
diag.primary_message(format!(
"token-program value `{identity}` was reassigned between token \
operations"
));
diag.help(
"validate the token-program account once, copy its address into an \
immutable binding, and reuse that binding for every token operation",
);
});
}
previous_usage = Some(call);
continue;
}
cx.lint(REQUIRE_CONSISTENT_TOKEN_PROGRAM, |diag| {
diag.span(call.span);
diag.primary_message(format!(
"token operation uses `{identity}` after the instruction established `{first}`"
));
diag.help(
"validate the token-program account once, copy its address, and pass that \
same value to token parsing, ATA checks, and CPI",
);
});
previous_usage = Some(call);
}
}
}