extern crate rustc_hir;
extern crate rustc_middle;
extern crate rustc_span;
use std::collections::HashMap;
use std::collections::HashSet;
use rustc_hir::BinOpKind;
use rustc_hir::Expr;
use rustc_hir::ExprKind;
use rustc_hir::HirId;
use rustc_hir::MatchSource;
use rustc_hir::Node;
use rustc_hir::Pat;
use rustc_hir::PatKind;
use rustc_hir::def::DefKind;
use rustc_hir::def::Res;
use rustc_hir::def_id::DefId;
use rustc_hir::intravisit::FnKind;
use rustc_hir::intravisit::Visitor;
use rustc_lint::LateContext;
use rustc_lint::LateLintPass;
use rustc_lint::LintContext;
use rustc_middle::ty::Ty;
use rustc_middle::ty::TyKind;
use rustc_span::Span;
use rustc_span::sym;
use crate::diagnostics;
use crate::shared;
crate::declare_late_lint! {
pub REQUIRE_POST_CPI_BALANCE_RELOAD,
Deny,
"token balances snapshotted before a value-moving CPI must be reloaded after it"
}
const TOKEN_PROGRAM_CRATES: &[&str] = &[
"pina",
"pinocchio_token",
"pinocchio_token_2022",
"spl_token",
"spl_token_2022",
"spl_token_interface",
];
const SYSTEM_PROGRAM_CRATES: &[&str] = &["pinocchio_system", "solana_system_interface"];
const TOKEN_VIEW_METHODS: &[&str] = &[
"as_account",
"as_associated_token_account",
"as_token_2022_account",
"as_token_account",
"as_token_account_for_program",
];
const TOKEN_VIEW_FUNCTIONS: &[&str] = &[
"from_account_info",
"from_account_info_unchecked",
"from_account_view",
"from_account_view_unchecked",
];
const UNWRAPPING_METHODS: &[&str] = &["expect", "map_err", "ok_or", "ok_or_else", "unwrap"];
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum TokenCpiKind {
Transfer,
TransferChecked,
MintTo,
MintToChecked,
}
impl TokenCpiKind {
fn from_type_name(name: &str) -> Option<Self> {
if name.ends_with("TransferChecked") {
Some(Self::TransferChecked)
} else if name.ends_with("MintToChecked") {
Some(Self::MintToChecked)
} else if name.ends_with("Transfer") {
Some(Self::Transfer)
} else if name.ends_with("MintTo") {
Some(Self::MintTo)
} else {
None
}
}
fn destination_index(self, leading_accounts: usize) -> usize {
if self.is_transfer() && leading_accounts >= 4 {
2
} else {
1
}
}
fn minimum_accounts(self, from_token_crate: bool) -> usize {
if self.is_transfer() && !from_token_crate {
4
} else {
3
}
}
fn is_transfer(self) -> bool {
matches!(self, Self::Transfer | Self::TransferChecked)
}
fn verb(self) -> &'static str {
if self.is_transfer() {
"transfer"
} else {
"mint"
}
}
}
struct TokenCpiConstructor {
kind: TokenCpiKind,
destination_index: usize,
display: String,
destination: Option<AccountKey>,
builder: DefId,
}
struct Invocation {
hir_id: HirId,
receiver: Option<DefId>,
receiver_is_legacy_builder: bool,
}
#[derive(Clone)]
struct AccountKey {
id: String,
text: String,
view: bool,
through_mut_binding: bool,
}
struct AmountRead {
hir_id: HirId,
span: Span,
account: Option<String>,
through_mut_binding: bool,
}
struct Definition<'tcx> {
bindings: Vec<HirId>,
span: Span,
value: &'tcx Expr<'tcx>,
}
struct LocalUse {
binding: HirId,
hir_id: HirId,
span: Span,
}
struct Analyzer<'cx, 'tcx> {
cx: &'cx LateContext<'tcx>,
let_initializers: HashMap<HirId, &'tcx Expr<'tcx>>,
mut_bindings: HashMap<HirId, &'tcx Expr<'tcx>>,
constructors: HashMap<Span, TokenCpiConstructor>,
invocations: HashMap<Span, Invocation>,
amount_reads: Vec<AmountRead>,
definitions: Vec<Definition<'tcx>>,
local_uses: Vec<LocalUse>,
closure_depth: usize,
}
fn local_path(expr: &Expr<'_>) -> Option<HirId> {
match expr.kind {
ExprKind::Path(rustc_hir::QPath::Resolved(None, path)) => {
match path.res {
Res::Local(binding) => Some(binding),
_ => None,
}
}
ExprKind::DropTemps(inner) => local_path(inner),
_ => None,
}
}
fn precedes(earlier: Span, later: Span) -> bool {
earlier.source_callsite().hi() <= later.source_callsite().lo()
}
fn encloses(outer: Span, inner: Span) -> bool {
outer.source_callsite().contains(inner.source_callsite())
}
fn node_span(node: Node<'_>) -> Option<Span> {
match node {
Node::Expr(expr) => Some(expr.span),
Node::Block(block) => Some(block.span),
_ => None,
}
}
fn is_cpi_method(method: &str) -> bool {
matches!(
method,
"invoke"
| "invoke_signed"
| "invoke_with_program"
| "invoke_signed_with_program"
| "invoke_with_unverified_program"
| "invoke_signed_with_unverified_program"
)
}
#[derive(Clone, Copy)]
struct ValueContext<'a> {
destination: &'a str,
invocation: Span,
target: Span,
}
fn is_comparison(operation: &str) -> bool {
matches!(
operation,
"eq" | "ne" | "lt" | "le" | "gt" | "ge" | "cmp" | "partial_cmp"
)
}
fn is_subtraction_like(operation: &str) -> bool {
matches!(
operation,
"sub" | "checked_sub" | "saturating_sub" | "wrapping_sub" | "overflowing_sub" | "abs_diff"
)
}
fn is_addition_like(operation: &str) -> bool {
matches!(
operation,
"add" | "checked_add" | "saturating_add" | "wrapping_add" | "overflowing_add"
)
}
impl<'tcx> Analyzer<'_, 'tcx> {
fn crate_name(&self, definition: DefId) -> String {
self.cx.tcx.crate_name(definition.krate).as_str().to_owned()
}
fn snippet(&self, span: Span) -> Option<String> {
self.cx.sess().source_map().span_to_snippet(span).ok()
}
fn core_adt(&self, mut ty: Ty<'tcx>) -> Option<DefId> {
loop {
ty = ty.peel_refs();
let TyKind::Adt(definition, generics) = ty.kind() else {
return None;
};
if self
.cx
.tcx
.is_diagnostic_item(sym::Result, definition.did())
|| self
.cx
.tcx
.is_diagnostic_item(sym::Option, definition.did())
{
ty = generics.type_at(0);
continue;
}
return Some(definition.did());
}
}
fn is_legacy_builder(&self, ty: Ty<'tcx>) -> bool {
let TyKind::Adt(builder, generics) = ty.peel_refs().kind() else {
return false;
};
self.crate_name(builder.did()) == "pinocchio_token"
&& generics.types().last().is_none_or(|program| {
program.ty_adt_def().is_some_and(|program| {
self.crate_name(program.did()) == "pinocchio_token"
&& self.cx.tcx.item_name(program.did()).as_str() == "TokenProgram"
})
})
}
fn is_constant(&self, expr: &Expr<'_>) -> bool {
match &expr.kind {
ExprKind::Lit(_) => true,
ExprKind::Tup([]) => true,
ExprKind::Unary(rustc_hir::UnOp::Neg, inner)
| ExprKind::AddrOf(_, _, inner)
| ExprKind::Cast(inner, _)
| ExprKind::DropTemps(inner) => self.is_constant(inner),
ExprKind::Path(path) => {
matches!(
self.cx.qpath_res(path, expr.hir_id),
Res::Def(
DefKind::Const
| DefKind::AssocConst | DefKind::ConstParam
| DefKind::Static { .. },
_
)
)
}
_ => false,
}
}
fn is_advancing_method(&self, call: &Expr<'_>) -> bool {
let Some(method) = self.cx.typeck_results().type_dependent_def_id(call.hir_id) else {
return false;
};
let name = self.cx.tcx.item_name(method);
let name = name.as_str();
if let Some(trait_id) = self.cx.tcx.trait_of_assoc(method) {
let trait_name = self.cx.tcx.item_name(trait_id);
return self.crate_name(trait_id) == "core"
&& matches!(trait_name.as_str(), "Iterator" | "DoubleEndedIterator")
&& matches!(name, "next" | "nth" | "next_back" | "nth_back");
}
self.crate_name(method) == "pina"
&& name.starts_with("next")
&& self
.cx
.tcx
.impl_of_assoc(method)
.and_then(|implementation| {
self.cx
.tcx
.type_of(implementation)
.instantiate_identity()
.skip_normalization()
.ty_adt_def()
})
.is_some_and(|cursor| {
self.cx.tcx.item_name(cursor.did()).as_str() == "AccountsCursor"
})
}
fn mutates_receiver(&self, initializer: &'tcx Expr<'tcx>) -> bool {
let call = self.peel_value(initializer);
let ExprKind::MethodCall(..) = call.kind else {
return false;
};
let Some(method) = self.cx.typeck_results().type_dependent_def_id(call.hir_id) else {
return false;
};
let signature = self
.cx
.tcx
.fn_sig(method)
.instantiate_identity()
.skip_binder();
signature.inputs().first().is_some_and(|receiver| {
matches!(
receiver.kind(),
TyKind::Ref(_, _, rustc_middle::ty::Mutability::Mut)
)
})
}
fn passes_receiver_through(&self, call: &Expr<'_>, method: &str) -> bool {
let Some(definition) = self.cx.typeck_results().type_dependent_def_id(call.hir_id) else {
return false;
};
match self.crate_name(definition).as_str() {
"core" => UNWRAPPING_METHODS.contains(&method),
"pina" => method.starts_with("assert_"),
_ => false,
}
}
fn is_token_view_function(&self, callee: &Expr<'_>) -> bool {
let ExprKind::Path(path) = &callee.kind else {
return false;
};
let Res::Def(DefKind::AssocFn, function) = self.cx.qpath_res(path, callee.hir_id) else {
return false;
};
TOKEN_VIEW_FUNCTIONS.contains(&self.cx.tcx.item_name(function).as_str())
&& TOKEN_PROGRAM_CRATES.contains(&self.crate_name(function).as_str())
}
fn account_key(&self, expr: &'tcx Expr<'tcx>) -> Option<AccountKey> {
match &expr.kind {
ExprKind::Path(rustc_hir::QPath::Resolved(None, path)) => {
let name = path
.segments
.iter()
.map(|segment| segment.ident.name.as_str())
.collect::<Vec<_>>()
.join("::");
let Res::Local(binding) = path.res else {
return Some(AccountKey {
id: name.clone(),
text: name,
view: false,
through_mut_binding: false,
});
};
if let Some(initializer) = self.let_initializers.get(&binding) {
return self.account_key(initializer);
}
if let Some(initializer) = self.mut_bindings.get(&binding) {
let call = self.account_key(initializer);
return Some(AccountKey {
id: match &call {
Some(call) => format!("{}~{name}", call.id),
None => format!("{name}#{}", binding.local_id.as_u32()),
},
text: name,
view: call.is_some_and(|call| call.view),
through_mut_binding: true,
});
}
Some(AccountKey {
id: format!("{name}#{}", binding.local_id.as_u32()),
text: name,
view: false,
through_mut_binding: false,
})
}
ExprKind::Field(base, field) => {
let base = self.account_key(base)?;
if base.view {
return (field.name.as_str() == "base").then_some(base);
}
Some(AccountKey {
id: format!("{}.{field}", base.id),
text: format!("{}.{field}", base.text),
view: false,
through_mut_binding: base.through_mut_binding,
})
}
ExprKind::Index(base, index, _) => {
let base = self.account_key(base)?;
if base.view || !matches!(index.kind, ExprKind::Lit(_)) {
return None;
}
let index = self.snippet(index.span)?;
Some(AccountKey {
id: format!("{}[{index}]", base.id),
text: format!("{}[{index}]", base.text),
view: false,
through_mut_binding: base.through_mut_binding,
})
}
ExprKind::Unary(rustc_hir::UnOp::Deref, inner)
| ExprKind::AddrOf(_, _, inner)
| ExprKind::DropTemps(inner)
| ExprKind::Type(inner, _)
| ExprKind::Match(inner, _, MatchSource::TryDesugar(_)) => self.account_key(inner),
ExprKind::Call(callee, arguments) => {
if let Some(argument) = shared::try_branch_argument(self.cx, expr) {
return self.account_key(argument);
}
let [account, ..] = arguments else {
return None;
};
if !self.is_token_view_function(callee) {
return None;
}
let account = self.account_key(account)?;
Some(AccountKey {
view: true,
..account
})
}
ExprKind::MethodCall(segment, receiver, arguments, _) => {
let receiver_key = self.account_key(receiver)?;
let method = segment.ident.name.as_str();
if TOKEN_VIEW_METHODS.contains(&method) && !receiver_key.view {
return Some(AccountKey {
view: true,
..receiver_key
});
}
if self.is_advancing_method(expr) {
return Some(AccountKey {
id: format!(
"{}.{method}@{}",
receiver_key.id,
expr.hir_id.local_id.as_u32()
),
text: format!("{}.{method}()", receiver_key.text),
view: false,
through_mut_binding: receiver_key.through_mut_binding,
});
}
if self.passes_receiver_through(expr, method) {
return Some(receiver_key);
}
if !arguments.iter().all(|argument| self.is_constant(argument)) {
return None;
}
let definition = self
.cx
.typeck_results()
.type_dependent_def_id(expr.hir_id)?;
let arguments = arguments
.iter()
.map(|argument| self.snippet(argument.span))
.collect::<Option<Vec<_>>>()?
.join(", ");
Some(AccountKey {
id: format!("{}.{definition:?}({arguments})", receiver_key.id),
text: format!("{}.{method}({arguments})", receiver_key.text),
view: false,
through_mut_binding: receiver_key.through_mut_binding,
})
}
ExprKind::Block(block, _) => block.expr.and_then(|tail| self.account_key(tail)),
_ => None,
}
}
fn token_cpi_constructor(
&self,
expr: &'tcx Expr<'tcx>,
callee: &'tcx Expr<'tcx>,
args: &'tcx [Expr<'tcx>],
) -> Option<TokenCpiConstructor> {
let ExprKind::Path(path) = &callee.kind else {
return None;
};
let Res::Def(DefKind::AssocFn, function) = self.cx.qpath_res(path, callee.hir_id) else {
return None;
};
if !matches!(
self.cx.tcx.item_name(function).as_str(),
"new" | "with_multisig_signers"
) {
return None;
}
let builder = self.core_adt(self.cx.typeck_results().expr_ty(expr))?;
let builder_crate = self.crate_name(builder);
if SYSTEM_PROGRAM_CRATES.contains(&builder_crate.as_str()) {
return None;
}
let kind = TokenCpiKind::from_type_name(self.cx.tcx.item_name(builder).as_str())?;
let signature = self
.cx
.tcx
.fn_sig(function)
.instantiate_identity()
.skip_binder();
let parameters = signature.inputs();
let leading_accounts = parameters
.iter()
.take_while(|parameter| {
matches!(parameter.kind(), TyKind::Ref(_, inner, _)
if matches!(inner.kind(), TyKind::Adt(..) | TyKind::Param(_)))
})
.count();
let has_amount = parameters
.get(leading_accounts)
.is_some_and(|amount| amount.is_integral());
let from_token_crate = TOKEN_PROGRAM_CRATES.contains(&builder_crate.as_str());
if leading_accounts < kind.minimum_accounts(from_token_crate) || !has_amount {
return None;
}
let destination_index = kind.destination_index(leading_accounts);
let destination = args.get(destination_index)?;
let display = self.snippet(destination.span)?;
let display = display
.trim_start_matches("&mut ")
.trim_start_matches('&')
.to_owned();
Some(TokenCpiConstructor {
kind,
destination_index,
display,
destination: self.account_key(destination),
builder,
})
}
fn record_definitions(&mut self, pattern: &Pat<'_>, value: &'tcx Expr<'tcx>, span: Span) {
if let (PatKind::Tuple(patterns, rest), ExprKind::Tup(values)) =
(&pattern.kind, &value.kind)
&& rest.as_opt_usize().is_none()
&& patterns.len() == values.len()
{
for (pattern, value) in patterns.iter().zip(values.iter()) {
self.record_definitions(pattern, value, span);
}
return;
}
let typeck = self.cx.typeck_results();
let mut bindings = Vec::new();
pattern.each_binding(|_, binding, _, _| {
if typeck.node_type(binding).is_integral() {
bindings.push(binding);
}
});
if bindings.is_empty() {
return;
}
self.definitions.push(Definition {
bindings,
span,
value,
});
}
fn reads_of<'a>(&'a self, destination: &'a str) -> impl Iterator<Item = &'a AmountRead> {
self.amount_reads
.iter()
.filter(move |read| read.account.as_deref() == Some(destination))
}
fn snapshot_definitions(&self, value: &ValueContext<'_>) -> Vec<&Definition<'tcx>> {
self.definitions
.iter()
.filter(|definition| precedes(definition.span, value.invocation))
.filter(|definition| self.is_snapshot_derived(definition.value, value, 0))
.collect()
}
fn reaching_definition(&self, binding: HirId, point: Span) -> Option<&Definition<'tcx>> {
self.definitions
.iter()
.filter(|definition| {
definition.bindings.contains(&binding) && precedes(definition.span, point)
})
.max_by_key(|definition| definition.span.source_callsite().lo())
}
fn diverges(&self, expr: &Expr<'_>) -> bool {
match expr.kind {
ExprKind::Ret(_)
| ExprKind::Break(..)
| ExprKind::Continue(_)
| ExprKind::Become(_) => true,
ExprKind::Block(..) | ExprKind::Loop(..) => {
self.cx.typeck_results().expr_ty(expr).is_never()
}
_ => false,
}
}
fn reaches(&self, invocation: HirId, target: Span) -> bool {
let mut child = invocation;
for (id, node) in self.cx.tcx.hir_parent_iter(invocation) {
if node_span(node).is_some_and(|span| encloses(span, target)) {
return match node {
Node::Expr(Expr {
kind: ExprKind::If(condition, ..),
..
}) => condition.hir_id == child,
Node::Expr(Expr {
kind: ExprKind::Match(_, arms, _),
..
}) => !arms.iter().any(|arm| arm.hir_id == child),
_ => true,
};
}
if let Node::Expr(expr) = node
&& self.diverges(expr)
{
return false;
}
child = id;
}
true
}
fn dominates(&self, read: HirId, target: Span) -> bool {
let mut child = read;
for (id, node) in self.cx.tcx.hir_parent_iter(read) {
let conditional = match node {
Node::Arm(_) => true,
Node::Expr(expr) => {
match expr.kind {
ExprKind::If(condition, ..) => condition.hir_id != child,
ExprKind::Loop(..) | ExprKind::Closure(..) => true,
ExprKind::Binary(operator, _, right) => {
matches!(operator.node, BinOpKind::And | BinOpKind::Or)
&& right.hir_id == child
}
_ => false,
}
}
_ => false,
};
if conditional {
return false;
}
if node_span(node).is_some_and(|span| encloses(span, target)) {
return true;
}
child = id;
}
false
}
fn peel_value(&self, mut expr: &'tcx Expr<'tcx>) -> &'tcx Expr<'tcx> {
loop {
if let Some(operand) = self.integer_conversion_operand(expr) {
expr = operand;
continue;
}
expr = match &expr.kind {
ExprKind::DropTemps(inner)
| ExprKind::Cast(inner, _)
| ExprKind::Type(inner, _)
| ExprKind::AddrOf(_, _, inner)
| ExprKind::Unary(rustc_hir::UnOp::Deref, inner)
| ExprKind::Match(inner, _, MatchSource::TryDesugar(_)) => inner,
ExprKind::Call(..) => {
match shared::try_branch_argument(self.cx, expr) {
Some(inner) => inner,
None => return expr,
}
}
ExprKind::MethodCall(segment, receiver, ..)
if UNWRAPPING_METHODS.contains(&segment.ident.name.as_str()) =>
{
receiver
}
_ => return expr,
};
}
}
fn is_conversion(&self, function: DefId) -> bool {
let tcx = self.cx.tcx;
let conversion_trait = tcx.trait_of_assoc(function).or_else(|| {
tcx.trait_impl_of_assoc(function)
.map(|implementation| tcx.impl_trait_id(implementation))
});
conversion_trait.is_some_and(|conversion| {
[sym::From, sym::Into, sym::TryFrom, sym::TryInto]
.into_iter()
.any(|name| tcx.is_diagnostic_item(name, conversion))
})
}
fn is_integer_payload(&self, ty: Ty<'tcx>) -> bool {
match ty.kind() {
TyKind::Adt(definition, generics)
if self
.cx
.tcx
.is_diagnostic_item(sym::Result, definition.did()) =>
{
generics.type_at(0).is_integral()
}
_ => ty.is_integral(),
}
}
fn integer_conversion_operand(&self, expr: &'tcx Expr<'tcx>) -> Option<&'tcx Expr<'tcx>> {
let typeck = self.cx.typeck_results();
let (function, operand) = match &expr.kind {
ExprKind::Call(callee, [operand]) => {
let ExprKind::Path(path) = &callee.kind else {
return None;
};
let Res::Def(DefKind::AssocFn, function) = self.cx.qpath_res(path, callee.hir_id)
else {
return None;
};
(function, operand)
}
ExprKind::MethodCall(_, receiver, [], _) => {
(typeck.type_dependent_def_id(expr.hir_id)?, *receiver)
}
_ => return None,
};
let converts_integers = self.is_conversion(function)
&& typeck.expr_ty(operand).peel_refs().is_integral()
&& self.is_integer_payload(typeck.expr_ty(expr));
converts_integers.then_some(operand)
}
fn arithmetic_parts(&self, expr: &'tcx Expr<'tcx>) -> Option<(String, Vec<&'tcx Expr<'tcx>>)> {
match &expr.kind {
ExprKind::Binary(operator, left, right) => {
let name = match operator.node {
BinOpKind::Add => "add",
BinOpKind::Sub => "sub",
BinOpKind::Mul => "mul",
BinOpKind::Div => "div",
BinOpKind::Rem => "rem",
BinOpKind::Eq => "eq",
BinOpKind::Ne => "ne",
BinOpKind::Lt => "lt",
BinOpKind::Le => "le",
BinOpKind::Gt => "gt",
BinOpKind::Ge => "ge",
BinOpKind::BitAnd => "bitand",
BinOpKind::BitOr => "bitor",
BinOpKind::BitXor => "bitxor",
BinOpKind::Shl => "shl",
BinOpKind::Shr => "shr",
BinOpKind::And | BinOpKind::Or => return None,
};
Some((name.to_owned(), vec![*left, *right]))
}
ExprKind::MethodCall(segment, receiver, arguments, _) => {
let integer_receiver = self
.cx
.typeck_results()
.expr_ty(receiver)
.peel_refs()
.is_integral();
integer_receiver.then(|| {
let mut operands = vec![*receiver];
operands.extend(arguments.iter());
(segment.ident.name.as_str().to_owned(), operands)
})
}
ExprKind::Call(callee, arguments) => {
let ExprKind::Path(path) = &callee.kind else {
return None;
};
let Res::Def(DefKind::AssocFn, function) = self.cx.qpath_res(path, callee.hir_id)
else {
return None;
};
let integer_first = arguments.first().is_some_and(|first| {
self.cx
.typeck_results()
.expr_ty(first)
.peel_refs()
.is_integral()
});
integer_first.then(|| {
(
self.cx.tcx.item_name(function).as_str().to_owned(),
arguments.iter().collect(),
)
})
}
_ => None,
}
}
fn local_definition(&self, expr: &'tcx Expr<'tcx>) -> Option<&Definition<'tcx>> {
let binding = local_path(expr)?;
self.reaching_definition(binding, expr.span)
}
fn is_reload(&self, expr: &'tcx Expr<'tcx>, value: &ValueContext<'_>, depth: usize) -> bool {
if depth > 16 {
return false;
}
let expr = self.peel_value(expr);
if self.amount_reads.iter().any(|read| {
read.hir_id == expr.hir_id
&& read.account.as_deref() == Some(value.destination)
&& precedes(value.invocation, read.span)
&& self.dominates(read.hir_id, value.target)
}) {
return true;
}
self.local_definition(expr).is_some_and(|definition| {
precedes(value.invocation, definition.span)
&& self.is_reload(definition.value, value, depth + 1)
})
}
fn is_delta(&self, expr: &'tcx Expr<'tcx>, value: &ValueContext<'_>, depth: usize) -> bool {
if depth > 16 {
return false;
}
let expr = self.peel_value(expr);
if let Some(definition) = self.local_definition(expr) {
return precedes(value.invocation, definition.span)
&& self.is_delta(definition.value, value, depth + 1);
}
let Some((operation, operands)) = self.arithmetic_parts(expr) else {
return false;
};
let [left, right] = operands[..] else {
return false;
};
let forward = self.is_reload(left, value, depth + 1)
&& self.is_snapshot_derived(right, value, depth + 1);
let symmetric = operation == "abs_diff"
&& self.is_snapshot_derived(left, value, depth + 1)
&& self.is_reload(right, value, depth + 1);
is_subtraction_like(&operation) && (forward || symmetric)
}
fn is_reload_derived(
&self,
expr: &'tcx Expr<'tcx>,
value: &ValueContext<'_>,
depth: usize,
) -> bool {
if depth > 16 {
return false;
}
if self.is_reload(expr, value, depth + 1) || self.is_delta(expr, value, depth + 1) {
return true;
}
let expr = self.peel_value(expr);
if let Some(definition) = self.local_definition(expr) {
return precedes(value.invocation, definition.span)
&& self.is_reload_derived(definition.value, value, depth + 1);
}
self.arithmetic_parts(expr).is_some_and(|(_, operands)| {
operands
.iter()
.any(|operand| self.is_reload_derived(operand, value, depth + 1))
})
}
fn is_snapshot_derived(
&self,
expr: &'tcx Expr<'tcx>,
value: &ValueContext<'_>,
depth: usize,
) -> bool {
if depth > 16 || self.is_delta(expr, value, depth + 1) {
return false;
}
let expr = self.peel_value(expr);
if self.amount_reads.iter().any(|read| {
read.hir_id == expr.hir_id
&& read.account.as_deref() == Some(value.destination)
&& precedes(read.span, value.invocation)
}) {
return true;
}
if let Some(definition) = self.local_definition(expr) {
return self.is_snapshot_derived(definition.value, value, depth + 1);
}
if let ExprKind::Block(block, _) = &expr.kind {
return block
.expr
.is_some_and(|tail| self.is_snapshot_derived(tail, value, depth + 1));
}
self.arithmetic_parts(expr).is_some_and(|(_, operands)| {
operands
.iter()
.any(|operand| self.is_snapshot_derived(operand, value, depth + 1))
})
}
fn stale_flow(
&self,
expr: &'tcx Expr<'tcx>,
value: &ValueContext<'_>,
depth: usize,
) -> Option<Span> {
if depth > 16 {
return Some(expr.span);
}
let mut current = expr;
loop {
let parent = self.cx.tcx.parent_hir_node(current.hir_id);
let Node::Expr(parent_expr) = parent else {
return self.stale_flow_into_binding(parent, current, value, depth);
};
let converts = self
.integer_conversion_operand(parent_expr)
.is_some_and(|operand| operand.hir_id == current.hir_id);
let passes_through = converts
|| match &parent_expr.kind {
ExprKind::DropTemps(_)
| ExprKind::Cast(..)
| ExprKind::Type(..)
| ExprKind::AddrOf(..)
| ExprKind::Unary(rustc_hir::UnOp::Deref, _) => true,
ExprKind::Match(scrutinee, _, MatchSource::TryDesugar(_)) => {
scrutinee.hir_id == current.hir_id
}
ExprKind::Call(..) => {
shared::try_branch_argument(self.cx, parent_expr)
.is_some_and(|argument| argument.hir_id == current.hir_id)
}
ExprKind::MethodCall(segment, receiver, ..) => {
receiver.hir_id == current.hir_id
&& UNWRAPPING_METHODS.contains(&segment.ident.name.as_str())
}
_ => false,
};
if passes_through {
current = parent_expr;
continue;
}
if let ExprKind::Assign(target, source, _) = &parent_expr.kind
&& source.hir_id == current.hir_id
&& local_path(target).is_some()
{
return self.stale_uses_of_definition(parent_expr.span, value, depth);
}
let Some((operation, operands)) = self.arithmetic_parts(parent_expr) else {
return Some(current.span);
};
let others = operands
.iter()
.filter(|operand| operand.hir_id != current.hir_id)
.copied()
.collect::<Vec<_>>();
if is_comparison(&operation) {
let verified = others.iter().all(|other| {
self.is_constant(other) || self.is_reload_derived(other, value, depth + 1)
});
return (!verified).then_some(current.span);
}
let is_subtrahend = operands.len() == 2 && operands[1].hir_id == current.hir_id;
if is_subtraction_like(&operation)
&& (is_subtrahend || operation == "abs_diff")
&& others
.iter()
.any(|other| self.is_reload(other, value, depth + 1))
{
return None;
}
if is_addition_like(&operation)
&& others
.iter()
.any(|other| self.is_delta(other, value, depth + 1))
{
return None;
}
current = parent_expr;
}
}
fn stale_flow_into_binding(
&self,
parent: Node<'tcx>,
current: &'tcx Expr<'tcx>,
value: &ValueContext<'_>,
depth: usize,
) -> Option<Span> {
let Node::LetStmt(local) = parent else {
return Some(current.span);
};
if local
.init
.is_none_or(|initializer| initializer.hir_id != current.hir_id)
{
return Some(current.span);
}
if matches!(local.pat.kind, PatKind::Wild) {
return None;
}
if !matches!(local.pat.kind, PatKind::Binding(_, _, _, None)) {
return Some(current.span);
}
self.stale_uses_of_definition(local.span, value, depth)
}
fn stale_uses_of_definition(
&self,
span: Span,
value: &ValueContext<'_>,
depth: usize,
) -> Option<Span> {
let Some(definition) = self
.definitions
.iter()
.find(|definition| definition.span == span)
else {
return Some(span);
};
self.local_uses
.iter()
.filter(|usage| definition.bindings.contains(&usage.binding))
.filter(|usage| {
self.reaching_definition(usage.binding, usage.span)
.is_some_and(|reaching| std::ptr::eq(reaching, definition))
})
.find_map(|usage| {
let use_expr = self.cx.tcx.hir_expect_expr(usage.hir_id);
let nested = ValueContext {
target: usage.span,
..*value
};
self.stale_flow(use_expr, &nested, depth + 1)
})
}
fn stale_snapshot_use(
&self,
destination: &str,
invocation: HirId,
invocation_span: Span,
) -> Option<Span> {
let context = ValueContext {
destination,
invocation: invocation_span,
target: invocation_span,
};
let snapshots = self.snapshot_definitions(&context);
self.local_uses
.iter()
.filter(|usage| precedes(invocation_span, usage.span))
.filter(|usage| {
self.reaching_definition(usage.binding, usage.span)
.is_some_and(|definition| {
snapshots
.iter()
.any(|snapshot| std::ptr::eq(*snapshot, definition))
})
})
.filter(|usage| self.reaches(invocation, usage.span))
.find_map(|usage| {
let use_expr = self.cx.tcx.hir_expect_expr(usage.hir_id);
let context = ValueContext {
target: usage.span,
..context
};
self.stale_flow(use_expr, &context, 0)
})
}
fn custody_transfer_is_unaccounted(
&self,
destination: &str,
invocation: Span,
cpi_spans: &[Span],
) -> bool {
let cpi_between = |start: Span, end: Span| {
cpi_spans
.iter()
.any(|cpi| precedes(start, *cpi) && precedes(*cpi, end))
};
let before = self
.reads_of(destination)
.filter(|read| precedes(read.span, invocation))
.max_by_key(|read| read.span.source_callsite().lo());
let after = self
.reads_of(destination)
.filter(|read| precedes(invocation, read.span))
.min_by_key(|read| read.span.source_callsite().lo());
let has_before = before.is_some_and(|read| !cpi_between(read.span, invocation));
let has_after = after.is_some_and(|read| !cpi_between(invocation, read.span));
!(has_before && has_after)
}
fn record_amount_read(&mut self, expr: &Expr<'_>, account: &'tcx Expr<'tcx>) {
if self.closure_depth > 0 {
return;
}
let key = self.account_key(account);
self.amount_reads.push(AmountRead {
hir_id: expr.hir_id,
span: expr.span,
through_mut_binding: key.as_ref().is_some_and(|key| key.through_mut_binding),
account: key.map(|key| key.id),
});
}
}
impl<'tcx> Visitor<'tcx> for Analyzer<'_, 'tcx> {
fn visit_local(&mut self, local: &'tcx rustc_hir::LetStmt<'tcx>) {
if let Some(initializer) = local.init {
if let PatKind::Binding(_, binding, _, None) = local.pat.kind
&& local.els.is_none()
{
if !self.mutates_receiver(initializer) {
self.let_initializers.insert(binding, initializer);
} else if !self.is_advancing_method(self.peel_value(initializer)) {
self.mut_bindings.insert(binding, initializer);
}
}
self.record_definitions(local.pat, initializer, local.span);
}
rustc_hir::intravisit::walk_local(self, local);
}
fn visit_expr(&mut self, expr: &'tcx Expr<'tcx>) {
match &expr.kind {
ExprKind::Call(callee, args) => {
if let Some(constructor) = self.token_cpi_constructor(expr, callee, args) {
self.constructors.insert(expr.span, constructor);
}
if let (ExprKind::Path(path), [account]) = (&callee.kind, args) {
let is_amount = match self.cx.qpath_res(path, callee.hir_id) {
Res::Def(DefKind::AssocFn | DefKind::Fn, function) => {
self.cx.tcx.item_name(function).as_str() == "amount"
}
_ => false,
};
if is_amount {
self.record_amount_read(expr, account);
}
}
}
ExprKind::MethodCall(segment, receiver, arguments, _) => {
let method = segment.ident.name.as_str();
if is_cpi_method(method) {
let receiver_ty = self.cx.typeck_results().expr_ty(receiver);
self.invocations.insert(
expr.span,
Invocation {
hir_id: expr.hir_id,
receiver: receiver_ty
.peel_refs()
.ty_adt_def()
.map(|definition| definition.did()),
receiver_is_legacy_builder: self.is_legacy_builder(receiver_ty),
},
);
}
if method == "amount" && arguments.is_empty() {
self.record_amount_read(expr, receiver);
}
}
ExprKind::Assign(target, value, _) => {
if let Some(binding) = local_path(target)
&& self.cx.typeck_results().expr_ty(target).is_integral()
{
self.definitions.push(Definition {
bindings: vec![binding],
span: expr.span,
value,
});
self.visit_expr(value);
return;
}
}
ExprKind::Path(rustc_hir::QPath::Resolved(None, path)) => {
if let Res::Local(binding) = path.res {
self.local_uses.push(LocalUse {
binding,
hir_id: expr.hir_id,
span: expr.span,
});
}
}
ExprKind::Closure(closure) => {
self.closure_depth += 1;
self.visit_body(self.cx.tcx.hir_body(closure.body));
self.closure_depth -= 1;
}
_ => {}
}
rustc_hir::intravisit::walk_expr(self, expr);
}
}
fn invocation_matches(constructor: &shared::CallInfo, invocation: &shared::CallInfo) -> bool {
if !is_cpi_method(&invocation.method) {
return false;
}
if let Some(binding) = constructor.result_binding.as_deref()
&& invocation.receiver.as_deref() == Some(binding)
{
return true;
}
invocation
.receiver_span
.is_some_and(|receiver| receiver.contains(constructor.span))
}
fn textual_custody_is_accounted(
calls: &[shared::CallInfo],
invocation: usize,
written: &str,
) -> bool {
let reads_amount = |candidate: &shared::CallInfo| {
candidate.method == "amount" && candidate.receiver.as_deref() == Some(written)
};
let has_before = calls[..invocation]
.iter()
.rposition(reads_amount)
.is_some_and(|before| {
!calls[before + 1..invocation]
.iter()
.any(|call| is_cpi_method(&call.method))
});
let has_after = calls[invocation + 1..]
.iter()
.position(reads_amount)
.is_some_and(|offset| {
!calls[invocation + 1..invocation + 1 + offset]
.iter()
.any(|call| is_cpi_method(&call.method))
});
has_before && has_after
}
fn is_custody_account(identity: &str) -> bool {
let name = identity.to_ascii_lowercase();
["vault", "custody", "reserve", "pool"]
.iter()
.any(|part| name.contains(part))
}
fn lint_custody_transfer(cx: &LateContext<'_>, invocation: Span, destination: &str) {
diagnostics::emit(cx, REQUIRE_POST_CPI_BALANCE_RELOAD, |diag| {
diag.span(invocation);
diag.primary_message(format!(
"transfer into `{destination}` is not accounted from its observed balance delta"
));
diag.help(
"read the destination amount before CPI, release the destination borrow (drop or \
scope it), invoke the transfer, reload the amount, and use `checked_sub` for the \
received value",
);
});
}
fn lint_stale_snapshot(
cx: &LateContext<'_>,
invocation: Span,
stale_use: Span,
kind: TokenCpiKind,
destination: &str,
) {
diagnostics::emit(cx, REQUIRE_POST_CPI_BALANCE_RELOAD, |diag| {
diag.span(invocation);
diag.primary_message(format!(
"{} into `{destination}` makes an earlier read of its balance stale",
kind.verb()
));
diag.span_note(stale_use, "the pre-CPI balance snapshot is used here");
diag.help(
"reload the destination amount after the CPI and account from \
`after.checked_sub(before)` instead of trusting the snapshot",
);
});
}
impl<'tcx> LateLintPass<'tcx> for RequirePostCpiBalanceReload {
fn check_fn(
&mut self,
cx: &LateContext<'tcx>,
_: FnKind<'tcx>,
_: &'tcx rustc_hir::FnDecl<'tcx>,
body: &'tcx rustc_hir::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, &["process", "instruction"])
{
return;
}
let facts = shared::collect_function_facts(cx, body);
let mut analyzer = Analyzer {
cx,
let_initializers: HashMap::new(),
mut_bindings: HashMap::new(),
constructors: HashMap::new(),
invocations: HashMap::new(),
amount_reads: Vec::new(),
definitions: Vec::new(),
local_uses: Vec::new(),
closure_depth: 0,
};
analyzer.visit_body(body);
analyzer
.definitions
.sort_by_key(|definition| definition.span.source_callsite().lo());
let cpi_spans = facts
.calls
.iter()
.filter(|call| is_cpi_method(&call.method))
.map(|call| call.span)
.collect::<Vec<_>>();
let unknown_read_spans = analyzer
.amount_reads
.iter()
.filter(|read| read.account.is_none())
.map(|read| read.span)
.collect::<HashSet<_>>();
for (index, call) in facts.calls.iter().enumerate() {
let Some(constructor) = analyzer.constructors.get(&call.span) else {
continue;
};
let Some(offset) = facts.calls[index + 1..]
.iter()
.position(|next| invocation_matches(call, next))
else {
continue;
};
let invocation_index = index + 1 + offset;
let invocation_call = &facts.calls[invocation_index];
let Some(invocation) = analyzer.invocations.get(&invocation_call.span) else {
continue;
};
let is_static_legacy_invoke =
matches!(invocation_call.method.as_str(), "invoke" | "invoke_signed")
&& invocation.receiver == Some(constructor.builder)
&& invocation.receiver_is_legacy_builder;
if is_static_legacy_invoke {
continue;
}
let written = call
.args
.get(constructor.destination_index)
.and_then(Option::as_deref);
let destination = constructor.destination.as_ref();
let is_custody = constructor.kind.is_transfer()
&& (is_custody_account(&constructor.display)
|| written.is_some_and(is_custody_account)
|| destination.is_some_and(|key| is_custody_account(&key.text)));
if is_custody {
let typed_unaccounted = destination.is_none_or(|key| {
analyzer.custody_transfer_is_unaccounted(
&key.id,
invocation_call.span,
&cpi_spans,
)
});
let identity_uncertain = destination.is_none_or(|key| key.through_mut_binding)
|| analyzer
.amount_reads
.iter()
.any(|read| read.through_mut_binding)
|| facts.calls.iter().any(|read| {
read.method == "amount"
&& read.receiver.as_deref() == written
&& unknown_read_spans.contains(&read.span)
});
let main_accepts = written.is_none_or(|written| {
!is_custody_account(written)
|| textual_custody_is_accounted(&facts.calls, invocation_index, written)
});
if typed_unaccounted && !(identity_uncertain && main_accepts) {
lint_custody_transfer(cx, invocation_call.span, &constructor.display);
continue;
}
}
let Some(destination) = destination else {
continue;
};
if let Some(stale_use) = analyzer.stale_snapshot_use(
&destination.id,
invocation.hir_id,
invocation_call.span,
) {
lint_stale_snapshot(
cx,
invocation_call.span,
stale_use,
constructor.kind,
&constructor.display,
);
}
}
}
}