#![allow(dead_code)]
use std::collections::{HashMap, HashSet};
use syn::spanned::Spanned;
use syn::visit::Visit;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Origin {
SlicePattern,
Index,
Iter,
}
#[derive(Clone, Debug, PartialEq, Eq, Hash)]
pub enum Validation {
Signer,
Owner,
Key,
KeyCompared,
Writable,
Uninitialized,
Discriminator,
LengthChecked,
}
#[derive(Clone, Debug, PartialEq, Eq, Hash)]
pub enum Use {
BorrowData { mut_: bool, unchecked: bool },
DeserializeState(String),
LamportsRead,
LamportsMut,
CpiAccount,
InvokeSignedSeeds,
AssignOwner,
CloseDrainLamports,
}
pub struct AccountBinding {
pub name: String,
pub origin: Origin,
pub validations: HashSet<Validation>,
pub uses: HashSet<Use>,
pub delegated: bool,
pub read_span: Option<proc_macro2::Span>,
pub authority_span: Option<proc_macro2::Span>,
}
impl AccountBinding {
fn new(name: String, origin: Origin) -> Self {
Self {
name,
origin,
validations: HashSet::new(),
uses: HashSet::new(),
delegated: false,
read_span: None,
authority_span: None,
}
}
pub fn reads_data(&self) -> bool {
self.uses
.iter()
.any(|u| matches!(u, Use::BorrowData { .. } | Use::DeserializeState(_)))
}
pub fn unchecked_read(&self) -> bool {
self.uses.iter().any(|u| {
matches!(
u,
Use::BorrowData {
unchecked: true,
..
}
)
})
}
}
pub struct Handler {
pub name: String,
pub bindings: Vec<AccountBinding>,
pub cpi_sites: Vec<CpiSite>,
}
pub struct CpiSite {
pub program_binding: Option<usize>,
pub span: proc_macro2::Span,
}
const ACCOUNT_SLICE_TYPES: &[&str] = &["AccountInfo", "AccountView"];
const LOADER_METHODS: &[&str] = &[
"load",
"load_mut",
"from_bytes",
"from_account_info",
"from_account_view",
"try_from_slice",
"unpack",
"deserialize",
];
const SYSVAR_TYPES: &[&str] = &[
"Rent",
"Clock",
"EpochSchedule",
"Fees",
"SlotHashes",
"StakeHistory",
"Instructions",
"RecentBlockhashes",
"EpochRewards",
"LastRestartSlot",
];
const CPI_FUNCS: &[&str] = &["invoke", "invoke_signed"];
const INSTRUCTION_TYPES: &[&str] = &["Instruction", "InstructionView"];
const LOG_MACROS: &[&str] = &[
"msg", "log", "sol_log", "println", "print", "eprintln", "eprint", "format", "write",
"writeln", "panic", "dbg", "emit",
];
fn validation_for_method(name: &str) -> Option<Validation> {
Some(match name {
"is_signer" => Validation::Signer,
"owner" => Validation::Owner,
"key" | "address" => Validation::Key,
"is_writable" => Validation::Writable,
"is_data_empty" | "data_is_empty" => Validation::Uninitialized,
"data_len" => Validation::LengthChecked,
_ => return None,
})
}
fn borrow_use_for_method(name: &str) -> Option<Use> {
Some(match name {
"try_borrow_data" | "try_borrow" => Use::BorrowData {
mut_: false,
unchecked: false,
},
"try_borrow_mut_data" | "try_borrow_mut" => Use::BorrowData {
mut_: true,
unchecked: false,
},
"borrow_data_unchecked" => Use::BorrowData {
mut_: false,
unchecked: true,
},
"borrow_mut_data_unchecked" => Use::BorrowData {
mut_: true,
unchecked: true,
},
_ => return None,
})
}
pub fn extract_handlers(file: &syn::File) -> Vec<Handler> {
let mut handlers = Vec::new();
collect_from_items(&file.items, &mut handlers);
handlers
}
fn collect_from_items(items: &[syn::Item], out: &mut Vec<Handler>) {
for item in items {
match item {
syn::Item::Fn(f) => {
if let Some(accounts_param) = accounts_param_name(&f.sig) {
out.push(extract_handler(
&f.sig.ident.to_string(),
accounts_param,
&f.block,
));
}
}
syn::Item::Mod(m) => {
if let Some((_, inner)) = &m.content {
collect_from_items(inner, out);
}
}
_ => {}
}
}
}
fn accounts_param_name(sig: &syn::Signature) -> Option<String> {
for input in &sig.inputs {
let syn::FnArg::Typed(pat_ty) = input else {
continue;
};
if !is_account_slice(&pat_ty.ty) {
continue;
}
if let syn::Pat::Ident(p) = pat_ty.pat.as_ref() {
return Some(p.ident.to_string());
}
}
None
}
fn is_account_slice(ty: &syn::Type) -> bool {
let syn::Type::Reference(r) = ty else {
return false;
};
let syn::Type::Slice(s) = r.elem.as_ref() else {
return false;
};
last_segment_ident(&s.elem).is_some_and(|id| ACCOUNT_SLICE_TYPES.contains(&id.as_str()))
}
fn extract_handler(name: &str, accounts_param: String, block: &syn::Block) -> Handler {
let mut ex = Extractor {
accounts_param,
bindings: Vec::new(),
index: HashMap::new(),
aliases: HashMap::new(),
macro_texts: Vec::new(),
instr_programs: HashMap::new(),
cpi_sites: Vec::new(),
};
ex.visit_block(block);
ex.apply_macro_validations();
Handler {
name: name.to_string(),
bindings: ex.bindings,
cpi_sites: ex.cpi_sites,
}
}
struct Extractor {
accounts_param: String,
bindings: Vec<AccountBinding>,
index: HashMap<String, usize>,
aliases: HashMap<String, String>,
macro_texts: Vec<(String, String)>,
instr_programs: HashMap<String, Option<usize>>,
cpi_sites: Vec<CpiSite>,
}
impl Extractor {
fn add_binding(&mut self, name: String, origin: Origin) {
if self.index.contains_key(&name) {
return;
}
self.index.insert(name.clone(), self.bindings.len());
self.bindings.push(AccountBinding::new(name, origin));
}
fn binding_idx(&self, name: &str) -> Option<usize> {
let canonical = self.aliases.get(name).map(String::as_str).unwrap_or(name);
self.index.get(canonical).copied()
}
fn resolve(&self, expr: &syn::Expr) -> Option<usize> {
match expr {
syn::Expr::Reference(r) => self.resolve(&r.expr),
syn::Expr::Paren(p) => self.resolve(&p.expr),
syn::Expr::Group(g) => self.resolve(&g.expr),
syn::Expr::Unary(syn::ExprUnary {
op: syn::UnOp::Deref(_),
expr,
..
}) => self.resolve(expr),
syn::Expr::Path(p) => {
let name = p.path.get_ident()?.to_string();
self.binding_idx(&name)
}
_ => None,
}
}
fn apply_macro_validations(&mut self) {
let texts = std::mem::take(&mut self.macro_texts);
for i in 0..self.bindings.len() {
let name = self.bindings[i].name.clone();
for (macro_name, text) in &texts {
if LOG_MACROS.contains(¯o_name.as_str()) {
continue;
}
for (method, val) in [
("owner", Validation::Owner),
("key", Validation::Key),
("address", Validation::Key),
("key", Validation::KeyCompared),
("address", Validation::KeyCompared),
("is_signer", Validation::Signer),
("is_writable", Validation::Writable),
("is_data_empty", Validation::Uninitialized),
("data_is_empty", Validation::Uninitialized),
("data_len", Validation::LengthChecked),
] {
if text.contains(&format!("{name} . {method}")) {
self.bindings[i].validations.insert(val.clone());
}
}
}
}
}
}
impl<'ast> Visit<'ast> for Extractor {
fn visit_local(&mut self, local: &'ast syn::Local) {
if let Some(init) = &local.init {
self.discover_binding(&local.pat, &init.expr);
if let syn::Pat::Ident(p) = &local.pat {
if let Some(s) = as_instruction_struct(&init.expr) {
let src = self.struct_program_binding(s);
self.instr_programs.insert(p.ident.to_string(), src);
}
}
}
syn::visit::visit_local(self, local);
}
fn visit_expr_binary(&mut self, node: &'ast syn::ExprBinary) {
if matches!(node.op, syn::BinOp::Eq(_) | syn::BinOp::Ne(_)) {
self.note_comparison(&node.left, &node.right, node.span());
}
syn::visit::visit_expr_binary(self, node);
}
fn visit_expr_method_call(&mut self, node: &'ast syn::ExprMethodCall) {
let method = node.method.to_string();
if method == "eq" || method == "ne" {
if let Some(arg) = node.args.first() {
self.note_comparison(&node.receiver, arg, node.span());
}
}
if let Some(idx) = self.resolve(&node.receiver) {
if let Some(v) = validation_for_method(&method) {
self.bindings[idx].validations.insert(v);
} else if let Some(u) = borrow_use_for_method(&method) {
if self.bindings[idx].read_span.is_none() {
self.bindings[idx].read_span = Some(node.method.span());
}
self.bindings[idx].uses.insert(u);
} else if method == "lamports" {
self.bindings[idx].uses.insert(Use::LamportsRead);
} else if method == "assign" {
self.bindings[idx].uses.insert(Use::AssignOwner);
}
}
syn::visit::visit_expr_method_call(self, node);
}
fn visit_expr_call(&mut self, node: &'ast syn::ExprCall) {
self.classify_call(node);
syn::visit::visit_expr_call(self, node);
}
fn visit_macro(&mut self, mac: &'ast syn::Macro) {
let name = mac
.path
.segments
.last()
.map(|s| s.ident.to_string())
.unwrap_or_default();
self.macro_texts.push((name, mac.tokens.to_string()));
syn::visit::visit_macro(self, mac);
}
}
impl Extractor {
fn discover_binding(&mut self, pat: &syn::Pat, init: &syn::Expr) {
match pat {
syn::Pat::Slice(slice) if self.is_accounts_expr(init) => {
for elem in &slice.elems {
if let syn::Pat::Ident(p) = elem {
self.add_binding(p.ident.to_string(), Origin::SlicePattern);
}
}
}
syn::Pat::Ident(p) => {
let name = p.ident.to_string();
if self.is_index_of_accounts(init) {
self.add_binding(name, Origin::Index);
} else if self.is_iter_next(init) {
self.add_binding(name, Origin::Iter);
} else if let Some(src) = self.resolve(init) {
let canonical = self.bindings[src].name.clone();
self.aliases.insert(name, canonical);
}
}
_ => {}
}
}
fn is_accounts_expr(&self, expr: &syn::Expr) -> bool {
matches!(expr, syn::Expr::Path(p) if p.path.is_ident(self.accounts_param.as_str()))
}
fn is_index_of_accounts(&self, expr: &syn::Expr) -> bool {
let expr = strip_ref(expr);
match expr {
syn::Expr::Index(idx) => self.is_accounts_expr(&idx.expr),
syn::Expr::MethodCall(m) => m.method == "get" && self.is_accounts_expr(&m.receiver),
_ => false,
}
}
fn is_iter_next(&self, expr: &syn::Expr) -> bool {
let expr = strip_try(expr);
matches!(expr, syn::Expr::Call(c)
if last_call_segment(&c.func).is_some_and(|s| s == "next_account_info"))
}
fn classify_call(&mut self, node: &syn::ExprCall) {
let Some((type_seg, fn_seg)) = call_path_tail(&node.func) else {
return;
};
if type_seg.is_none() && CPI_FUNCS.contains(&fn_seg.as_str()) {
for arg in &node.args {
if let Some(idx) = self.resolve(arg) {
self.bindings[idx].uses.insert(Use::CpiAccount);
}
}
let program_binding = node.args.first().and_then(|a| self.cpi_program_binding(a));
self.cpi_sites.push(CpiSite {
program_binding,
span: node.span(),
});
return;
}
if let Some(ty) = &type_seg {
if LOADER_METHODS.contains(&fn_seg.as_str()) && !SYSVAR_TYPES.contains(&ty.as_str()) {
for arg in &node.args {
if let Some(idx) = self.resolve(arg) {
if self.bindings[idx].read_span.is_none() {
self.bindings[idx].read_span = Some(node.span());
}
self.bindings[idx]
.uses
.insert(Use::DeserializeState(ty.clone()));
}
}
return;
}
}
for arg in &node.args {
if let Some(idx) = self.resolve(arg) {
self.bindings[idx].delegated = true;
}
}
}
fn note_comparison(&mut self, a: &syn::Expr, b: &syn::Expr, span: proc_macro2::Span) {
for (side, other) in [(a, b), (b, a)] {
if let Some(idx) = self.key_call_binding(side) {
self.bindings[idx]
.validations
.insert(Validation::KeyCompared);
if self.bindings[idx].authority_span.is_none() && is_authority_ref(other) {
self.bindings[idx].authority_span = Some(span);
}
}
}
}
fn key_call_binding(&self, expr: &syn::Expr) -> Option<usize> {
if let syn::Expr::MethodCall(m) = peel(expr) {
let method = m.method.to_string();
if method == "key" || method == "address" {
return self.resolve(&m.receiver);
}
}
None
}
fn cpi_program_binding(&self, arg0: &syn::Expr) -> Option<usize> {
if let Some(s) = as_instruction_struct(arg0) {
return self.struct_program_binding(s);
}
if let syn::Expr::Path(p) = peel(arg0) {
if let Some(id) = p.path.get_ident() {
return self.instr_programs.get(&id.to_string()).copied().flatten();
}
}
None
}
fn struct_program_binding(&self, s: &syn::ExprStruct) -> Option<usize> {
for fv in &s.fields {
if let syn::Member::Named(id) = &fv.member {
if id == "program_id" {
return self.key_call_binding(&fv.expr);
}
}
}
None
}
}
fn strip_ref(expr: &syn::Expr) -> &syn::Expr {
match expr {
syn::Expr::Reference(r) => strip_ref(&r.expr),
syn::Expr::Paren(p) => strip_ref(&p.expr),
_ => expr,
}
}
fn peel(expr: &syn::Expr) -> &syn::Expr {
match expr {
syn::Expr::Reference(r) => peel(&r.expr),
syn::Expr::Paren(p) => peel(&p.expr),
syn::Expr::Group(g) => peel(&g.expr),
syn::Expr::Unary(syn::ExprUnary {
op: syn::UnOp::Deref(_),
expr,
..
}) => peel(expr),
_ => expr,
}
}
fn is_authority_ref(expr: &syn::Expr) -> bool {
match peel(expr) {
syn::Expr::Field(f) => {
matches!(&f.member, syn::Member::Named(id) if is_authority_name(&id.to_string()))
}
syn::Expr::MethodCall(m) => is_authority_name(&m.method.to_string()),
syn::Expr::Path(p) => p
.path
.segments
.last()
.is_some_and(|s| is_authority_name(&s.ident.to_string())),
_ => false,
}
}
fn is_authority_name(name: &str) -> bool {
name == "authority"
|| name == "admin"
|| name == "auth"
|| name.ends_with("_authority")
|| name.ends_with("_auth")
}
fn as_instruction_struct(expr: &syn::Expr) -> Option<&syn::ExprStruct> {
if let syn::Expr::Struct(s) = peel(expr) {
let last = s.path.segments.last()?;
if INSTRUCTION_TYPES.contains(&last.ident.to_string().as_str()) {
return Some(s);
}
}
None
}
fn strip_try(expr: &syn::Expr) -> &syn::Expr {
match expr {
syn::Expr::Try(t) => strip_try(&t.expr),
syn::Expr::Paren(p) => strip_try(&p.expr),
_ => expr,
}
}
fn last_segment_ident(ty: &syn::Type) -> Option<String> {
if let syn::Type::Path(p) = ty {
return p.path.segments.last().map(|s| s.ident.to_string());
}
None
}
fn last_call_segment(func: &syn::Expr) -> Option<String> {
if let syn::Expr::Path(p) = func {
return p.path.segments.last().map(|s| s.ident.to_string());
}
None
}
fn call_path_tail(func: &syn::Expr) -> Option<(Option<String>, String)> {
let syn::Expr::Path(p) = func else {
return None;
};
let segs = &p.path.segments;
let fn_seg = segs.last()?.ident.to_string();
let type_seg = if segs.len() >= 2 {
Some(segs[segs.len() - 2].ident.to_string())
} else {
None
};
Some((type_seg, fn_seg))
}