Skip to main content

frs/cmds/rsl/rules/common/
fn_call_analysis.rs

1use std::collections::HashSet;
2
3use proc_macro2::Span;
4use syn::Expr;
5use syn::spanned::Spanned;
6use syn::visit::Visit;
7
8use super::fn_path_resolution::FnCallSuggestion;
9use super::fn_path_resolution::expected_fn_path;
10use super::module_idx::ModuleIdx;
11use super::module_idx::module_idx;
12use super::path_resolution::path_label;
13use super::path_resolution::path_parts;
14use super::scope_bindings::add_pattern_bindings;
15use super::scope_bindings::block_fn_bindings;
16use super::scope_bindings::closure_bindings;
17use super::scope_bindings::parameter_bindings;
18use super::scope_bindings::pattern_bindings;
19
20#[derive(Debug)]
21#[cfg_attr(test, derive(Eq, PartialEq))]
22pub struct CallDetails {
23    pub actual_path: String,
24    pub replacement_path: String,
25    pub add_import: Option<String>,
26}
27
28#[derive(Clone, Copy)]
29pub enum FnCallKind {
30    Unqualified,
31    Overqualified,
32}
33
34pub struct FnCallFinding {
35    pub span: Span,
36    pub actual_path: String,
37    pub suggestion: FnCallSuggestion,
38}
39
40pub fn find_fn_calls(file: &syn::File, kind: FnCallKind) -> Vec<FnCallFinding> {
41    let idx = module_idx(file);
42    let mut findings = Vec::new();
43
44    for scope in &idx.scopes {
45        let mut visitor = FnCallVisitor {
46            idx: &idx,
47            current_module: &scope.path,
48            kind,
49            findings: &mut findings,
50            local_bindings: Vec::new(),
51        };
52        for item in scope.items {
53            visitor.visit_item(item);
54        }
55    }
56
57    findings
58}
59
60struct FnCallVisitor<'idx, 'ast, 'output> {
61    idx: &'idx ModuleIdx<'ast>,
62    current_module: &'idx [String],
63    kind: FnCallKind,
64    findings: &'output mut Vec<FnCallFinding>,
65    local_bindings: Vec<HashSet<String>>,
66}
67
68impl<'ast> Visit<'ast> for FnCallVisitor<'_, '_, '_> {
69    fn visit_expr_call(&mut self, expression: &'ast syn::ExprCall) {
70        if let Expr::Path(path) = expression.func.as_ref()
71            && let Some(parts) = path_parts(&path.path)
72            && let Some(suggestion) = expected_fn_path(self.idx, self.current_module, &self.local_bindings, &path.path)
73        {
74            let is_relevant = match self.kind {
75                FnCallKind::Unqualified => parts.len() == 1,
76                FnCallKind::Overqualified => parts.len() > 1,
77            };
78            let actual_path = path_label(&path.path);
79            if is_relevant && actual_path != suggestion.expected_path {
80                self.findings.push(FnCallFinding {
81                    span: path.span(),
82                    actual_path,
83                    suggestion,
84                });
85            }
86        }
87
88        syn::visit::visit_expr_call(self, expression);
89    }
90
91    fn visit_block(&mut self, block: &'ast syn::Block) {
92        self.local_bindings.push(block_fn_bindings(block));
93        syn::visit::visit_block(self, block);
94        let _ = self.local_bindings.pop();
95    }
96
97    fn visit_expr_closure(&mut self, closure: &'ast syn::ExprClosure) {
98        self.local_bindings.push(closure_bindings(&closure.inputs));
99        syn::visit::visit_expr_closure(self, closure);
100        let _ = self.local_bindings.pop();
101    }
102
103    fn visit_expr_for_loop(&mut self, expression: &'ast syn::ExprForLoop) {
104        for attribute in &expression.attrs {
105            self.visit_attribute(attribute);
106        }
107        if let Some(label) = &expression.label {
108            self.visit_label(label);
109        }
110        self.visit_pat(&expression.pat);
111        self.visit_expr(&expression.expr);
112
113        self.local_bindings.push(pattern_bindings(&expression.pat));
114        self.visit_block(&expression.body);
115        let _ = self.local_bindings.pop();
116    }
117
118    fn visit_expr_if(&mut self, expression: &'ast syn::ExprIf) {
119        if let Expr::Let(let_expression) = expression.cond.as_ref() {
120            for attribute in &expression.attrs {
121                self.visit_attribute(attribute);
122            }
123            self.visit_expr(&expression.cond);
124            self.local_bindings.push(pattern_bindings(&let_expression.pat));
125            self.visit_block(&expression.then_branch);
126            let _ = self.local_bindings.pop();
127            if let Some((_, else_branch)) = &expression.else_branch {
128                self.visit_expr(else_branch);
129            }
130            return;
131        }
132
133        // TODO: Carry bindings through let chains in compound conditions.
134        syn::visit::visit_expr_if(self, expression);
135    }
136
137    fn visit_expr_while(&mut self, expression: &'ast syn::ExprWhile) {
138        if let Expr::Let(let_expression) = expression.cond.as_ref() {
139            for attribute in &expression.attrs {
140                self.visit_attribute(attribute);
141            }
142            if let Some(label) = &expression.label {
143                self.visit_label(label);
144            }
145            self.visit_expr(&expression.cond);
146            self.local_bindings.push(pattern_bindings(&let_expression.pat));
147            self.visit_block(&expression.body);
148            let _ = self.local_bindings.pop();
149            return;
150        }
151
152        // TODO: Carry bindings through let chains in compound conditions.
153        syn::visit::visit_expr_while(self, expression);
154    }
155
156    fn visit_impl_item_fn(&mut self, fn_item: &'ast syn::ImplItemFn) {
157        self.local_bindings.push(parameter_bindings(&fn_item.sig.inputs));
158        syn::visit::visit_impl_item_fn(self, fn_item);
159        let _ = self.local_bindings.pop();
160    }
161
162    fn visit_item_fn(&mut self, fn_item: &'ast syn::ItemFn) {
163        self.local_bindings.push(parameter_bindings(&fn_item.sig.inputs));
164        syn::visit::visit_item_fn(self, fn_item);
165        let _ = self.local_bindings.pop();
166    }
167
168    fn visit_local(&mut self, local: &'ast syn::Local) {
169        for attribute in &local.attrs {
170            self.visit_attribute(attribute);
171        }
172        self.visit_pat(&local.pat);
173        if let Some(init) = &local.init {
174            self.visit_expr(&init.expr);
175            if let Some((_, diverge)) = &init.diverge {
176                self.visit_expr(diverge);
177            }
178        }
179        if let Some(scope) = self.local_bindings.last_mut() {
180            add_pattern_bindings(scope, &local.pat);
181        }
182    }
183
184    fn visit_arm(&mut self, arm: &'ast syn::Arm) {
185        self.local_bindings.push(pattern_bindings(&arm.pat));
186        syn::visit::visit_arm(self, arm);
187        let _ = self.local_bindings.pop();
188    }
189
190    fn visit_trait_item_fn(&mut self, fn_item: &'ast syn::TraitItemFn) {
191        self.local_bindings.push(parameter_bindings(&fn_item.sig.inputs));
192        syn::visit::visit_trait_item_fn(self, fn_item);
193        let _ = self.local_bindings.pop();
194    }
195
196    fn visit_item_mod(&mut self, _module: &'ast syn::ItemMod) {}
197
198    fn visit_attribute(&mut self, _attribute: &'ast syn::Attribute) {}
199
200    fn visit_macro(&mut self, _mac: &'ast syn::Macro) {
201        // TODO: Resolve local bindings generated by macros.
202    }
203}