Skip to main content

frs/cmds/rsl/rules/
misordered_fn.rs

1//! Misordered-fn rule for `frs rsl`.
2
3use std::collections::HashMap;
4use std::collections::HashSet;
5use std::path::Path;
6
7use proc_macro2::Span;
8use syn::Expr;
9use syn::Item;
10use syn::visit::Visit;
11
12use super::common::Location;
13use crate::cmds::rsl::ast::ItemKind;
14use crate::cmds::rsl::ast::ModuleItem;
15use crate::cmds::rsl::ast::VisibilityClass;
16use crate::cmds::rsl::engine::FileContext;
17use crate::cmds::rsl::rules::TypedRule;
18use crate::cmds::rsl::rules::TypedRuleViolation;
19
20pub struct MisorderedFnRule;
21
22impl TypedRule for MisorderedFnRule {
23    type Violation = MisorderedFnViolation;
24
25    fn code() -> &'static str {
26        "misordered_fn"
27    }
28
29    fn check(&self, ctx: &FileContext<'_>) -> Vec<Self::Violation> {
30        let mut violations = Vec::new();
31
32        for items in &ctx.module_item_lists {
33            self::check_misordered_fn(&self::module_fns(items), ctx.path, &mut violations);
34
35            for module_item in items.iter().rev() {
36                let item = module_item.item();
37                if let Item::Impl(item_impl) = item
38                    && item_impl.trait_.is_none()
39                {
40                    self::check_misordered_fn(&self::impl_fns(item_impl), ctx.path, &mut violations);
41                }
42            }
43        }
44
45        violations
46    }
47}
48
49#[derive(Debug)]
50#[cfg_attr(test, derive(Eq, PartialEq))]
51pub struct MisorderedFnViolation {
52    pub location: Location,
53    pub details: MisorderedFnDetails,
54}
55
56impl MisorderedFnViolation {
57    fn new(path: &Path, span: Span, expected_after: String, item: ItemKind) -> Self {
58        Self {
59            location: Location::from_span(path, span),
60            details: MisorderedFnDetails { expected_after, item },
61        }
62    }
63}
64
65impl TypedRuleViolation for MisorderedFnViolation {
66    type Rule = MisorderedFnRule;
67}
68
69#[derive(Debug)]
70#[cfg_attr(test, derive(Eq, PartialEq))]
71pub struct MisorderedFnDetails {
72    pub expected_after: String,
73    pub item: ItemKind,
74}
75
76#[derive(Clone, Debug)]
77struct FnInfo {
78    source_idx: usize,
79    span: Span,
80    name: String,
81    visibility: VisibilityClass,
82    calls: Vec<String>,
83}
84
85#[derive(Default)]
86struct DirectCallCollector {
87    associated: bool,
88    calls: Vec<String>,
89}
90
91impl<'ast> Visit<'ast> for DirectCallCollector {
92    fn visit_expr_call(&mut self, expression: &'ast syn::ExprCall) {
93        if let Expr::Path(path) = expression.func.as_ref()
94            && let Some(name) = self::direct_call_name(path, self.associated)
95        {
96            self.calls.push(name);
97        }
98        syn::visit::visit_expr_call(self, expression);
99    }
100
101    fn visit_item_fn(&mut self, _fn: &'ast syn::ItemFn) {}
102}
103
104fn check_misordered_fn(fns: &[FnInfo], path: &Path, violations: &mut Vec<MisorderedFnViolation>) {
105    if fns.is_empty() {
106        return;
107    }
108
109    let targets = self::fn_targets(fns);
110    let callers_by_target = self::callers_by_target(fns, &targets);
111    let helper_components = self::helper_components(fns, &targets);
112    let mut component_for_helper = vec![None; fns.len()];
113    for (component_idx, component) in helper_components.iter().enumerate() {
114        for &helper in component {
115            if let Some(component_slot) = component_for_helper.get_mut(helper) {
116                *component_slot = Some(component_idx);
117            }
118        }
119    }
120
121    self::check_recursive_helpers(fns, &targets, &callers_by_target, &helper_components, path, violations);
122    self::check_non_recursive_helpers(
123        fns,
124        &targets,
125        &callers_by_target,
126        &helper_components,
127        &component_for_helper,
128        path,
129        violations,
130    );
131}
132
133fn fn_targets(fns: &[FnInfo]) -> HashMap<String, usize> {
134    let mut targets = HashMap::new();
135    let mut ambiguous = HashSet::new();
136    for (idx, fn_info) in fns.iter().enumerate() {
137        if targets.insert(fn_info.name.clone(), idx).is_some() {
138            ambiguous.insert(fn_info.name.clone());
139        }
140    }
141    for name in ambiguous {
142        targets.remove(&name);
143    }
144    targets
145}
146
147fn callers_by_target(fns: &[FnInfo], targets: &HashMap<String, usize>) -> Vec<Vec<usize>> {
148    let mut callers_by_target = vec![Vec::new(); fns.len()];
149    for (caller, fn_info) in fns.iter().enumerate() {
150        for called_name in &fn_info.calls {
151            let Some(&target) = targets.get(called_name) else {
152                continue;
153            };
154            let Some(target_fn) = fns.get(target) else {
155                continue;
156            };
157            let Some(target_callers) = callers_by_target.get_mut(target) else {
158                continue;
159            };
160            if target_fn.visibility != VisibilityClass::Private || target_callers.contains(&caller) {
161                continue;
162            }
163            target_callers.push(caller);
164        }
165    }
166    callers_by_target
167}
168
169fn check_recursive_helpers(
170    fns: &[FnInfo],
171    targets: &HashMap<String, usize>,
172    callers_by_target: &[Vec<usize>],
173    helper_components: &[Vec<usize>],
174    path: &Path,
175    violations: &mut Vec<MisorderedFnViolation>,
176) {
177    for component in helper_components {
178        if !self::component_is_recursive(component, fns, targets) {
179            continue;
180        }
181        self::check_recursive_anchor(component, fns, callers_by_target, path, violations);
182        self::check_recursive_adjacency(component, fns, path, violations);
183    }
184}
185
186fn check_recursive_anchor(
187    component: &[usize],
188    fns: &[FnInfo],
189    callers_by_target: &[Vec<usize>],
190    path: &Path,
191    violations: &mut Vec<MisorderedFnViolation>,
192) {
193    let members: HashSet<usize> = component.iter().copied().collect();
194    let external_callers: Vec<usize> = component
195        .iter()
196        .flat_map(|&helper| callers_by_target.get(helper).into_iter().flatten().copied())
197        .filter(|caller| !members.contains(caller))
198        .collect();
199    let Some(anchor) = self::caller_anchor(fns, &external_callers, true) else {
200        return;
201    };
202    let Some(&first) = component
203        .iter()
204        .min_by_key(|&&helper| fns.get(helper).map_or(usize::MAX, |fn_info| fn_info.source_idx))
205    else {
206        return;
207    };
208    let (Some(first_fn), Some(anchor_fn)) = (fns.get(first), fns.get(anchor)) else {
209        return;
210    };
211    if first_fn.source_idx < anchor_fn.source_idx && anchor_fn.visibility == VisibilityClass::Private {
212        self::push_caller_violation(first_fn, anchor_fn, path, violations);
213    }
214}
215
216fn check_recursive_adjacency(
217    component: &[usize],
218    fns: &[FnInfo],
219    path: &Path,
220    violations: &mut Vec<MisorderedFnViolation>,
221) {
222    let mut ordered = component.to_vec();
223    ordered.sort_unstable_by_key(|&helper| fns.get(helper).map_or(usize::MAX, |fn_info| fn_info.source_idx));
224    for pair in ordered.windows(2) {
225        let [previous, current] = pair else {
226            continue;
227        };
228        let (Some(previous_fn), Some(current_fn)) = (fns.get(*previous), fns.get(*current)) else {
229            continue;
230        };
231        if current_fn.source_idx != previous_fn.source_idx.saturating_add(1) {
232            self::push_caller_violation(current_fn, previous_fn, path, violations);
233        }
234    }
235}
236
237fn check_non_recursive_helpers(
238    fns: &[FnInfo],
239    targets: &HashMap<String, usize>,
240    callers_by_target: &[Vec<usize>],
241    helper_components: &[Vec<usize>],
242    component_for_helper: &[Option<usize>],
243    path: &Path,
244    violations: &mut Vec<MisorderedFnViolation>,
245) {
246    for (helper, callers) in callers_by_target.iter().enumerate() {
247        let Some(helper_fn) = fns.get(helper) else {
248            continue;
249        };
250        let recursive = component_for_helper
251            .get(helper)
252            .copied()
253            .flatten()
254            .and_then(|component| helper_components.get(component))
255            .is_some_and(|component| self::component_is_recursive(component, fns, targets));
256        if helper_fn.visibility != VisibilityClass::Private || recursive {
257            continue;
258        }
259        let Some(anchor) = self::caller_anchor(fns, callers, callers.len() > 1) else {
260            continue;
261        };
262        let Some(anchor_fn) = fns.get(anchor) else {
263            continue;
264        };
265        if helper_fn.source_idx < anchor_fn.source_idx && anchor_fn.visibility == VisibilityClass::Private {
266            self::push_caller_violation(helper_fn, anchor_fn, path, violations);
267        }
268    }
269}
270
271fn push_caller_violation(
272    item: &FnInfo,
273    expected_after: &FnInfo,
274    path: &Path,
275    violations: &mut Vec<MisorderedFnViolation>,
276) {
277    violations.push(MisorderedFnViolation::new(
278        path,
279        item.span,
280        format!("fn {}", expected_after.name),
281        ItemKind::Fn,
282    ));
283}
284
285fn caller_anchor(fns: &[FnInfo], callers: &[usize], prefer_public: bool) -> Option<usize> {
286    let public_callers = callers.iter().copied().filter(|&caller| {
287        fns.get(caller)
288            .is_some_and(|fn_info| fn_info.visibility == VisibilityClass::Public)
289    });
290    let candidates = if prefer_public {
291        let public_callers: Vec<_> = public_callers.collect();
292        if public_callers.is_empty() {
293            callers.to_vec()
294        } else {
295            public_callers
296        }
297    } else {
298        callers.to_vec()
299    };
300
301    candidates
302        .into_iter()
303        .min_by_key(|&caller| fns.get(caller).map_or(usize::MAX, |fn_info| fn_info.source_idx))
304}
305
306fn helper_components(fns: &[FnInfo], targets: &HashMap<String, usize>) -> Vec<Vec<usize>> {
307    let mut graph = vec![Vec::new(); fns.len()];
308    let mut reverse = vec![Vec::new(); fns.len()];
309    for (caller, fn_info) in fns.iter().enumerate() {
310        if fn_info.visibility != VisibilityClass::Private {
311            continue;
312        }
313        for called_name in &fn_info.calls {
314            let Some(&target) = targets.get(called_name) else {
315                continue;
316            };
317            let Some(target_fn) = fns.get(target) else {
318                continue;
319            };
320            let Some(caller_edges) = graph.get_mut(caller) else {
321                continue;
322            };
323            if target_fn.visibility != VisibilityClass::Private || caller_edges.contains(&target) {
324                continue;
325            }
326            caller_edges.push(target);
327            if let Some(target_edges) = reverse.get_mut(target) {
328                target_edges.push(caller);
329            }
330        }
331    }
332
333    let mut visited = vec![false; fns.len()];
334    let mut order = Vec::new();
335    for node in 0..fns.len() {
336        self::visit_graph(node, &graph, &mut visited, &mut order);
337    }
338
339    visited.fill(false);
340    let mut components = Vec::new();
341    for node in order.into_iter().rev() {
342        if visited.get(node).copied().unwrap_or(false) {
343            continue;
344        }
345        let mut component = Vec::new();
346        self::collect_graph_component(node, &reverse, &mut visited, &mut component);
347        components.push(component);
348    }
349
350    components
351}
352
353fn component_is_recursive(component: &[usize], fns: &[FnInfo], targets: &HashMap<String, usize>) -> bool {
354    if component.len() > 1 {
355        return true;
356    }
357    let Some(&helper) = component.first() else {
358        return false;
359    };
360    let Some(fn_info) = fns.get(helper) else {
361        return false;
362    };
363    fn_info
364        .calls
365        .iter()
366        .filter_map(|called_name| targets.get(called_name))
367        .any(|&target| target == helper)
368}
369
370fn visit_graph(node: usize, graph: &[Vec<usize>], visited: &mut [bool], order: &mut Vec<usize>) {
371    let mut stack = vec![(node, false)];
372    while let Some((current, expanded)) = stack.pop() {
373        if expanded {
374            order.push(current);
375            continue;
376        }
377        if visited.get(current).copied().unwrap_or(false) {
378            continue;
379        }
380        let Some(visited_node) = visited.get_mut(current) else {
381            continue;
382        };
383        *visited_node = true;
384        stack.push((current, true));
385        if let Some(next_nodes) = graph.get(current) {
386            for &next in next_nodes.iter().rev() {
387                stack.push((next, false));
388            }
389        }
390    }
391}
392
393fn collect_graph_component(node: usize, graph: &[Vec<usize>], visited: &mut [bool], component: &mut Vec<usize>) {
394    let mut stack = vec![node];
395    while let Some(current) = stack.pop() {
396        if visited.get(current).copied().unwrap_or(false) {
397            continue;
398        }
399        let Some(visited_node) = visited.get_mut(current) else {
400            continue;
401        };
402        *visited_node = true;
403        component.push(current);
404        if let Some(next_nodes) = graph.get(current) {
405            stack.extend(next_nodes.iter().rev().copied());
406        }
407    }
408}
409
410fn module_fns(items: &[ModuleItem<'_>]) -> Vec<FnInfo> {
411    items
412        .iter()
413        .enumerate()
414        .filter_map(|(source_idx, module_item)| {
415            let Item::Fn(fn_item) = module_item.item() else {
416                return None;
417            };
418            Some(FnInfo {
419                source_idx,
420                span: fn_item.sig.fn_token.span,
421                name: fn_item.sig.ident.to_string(),
422                visibility: VisibilityClass::from(&fn_item.vis),
423                calls: self::direct_calls(&fn_item.block, false),
424            })
425        })
426        .collect()
427}
428
429fn impl_fns(item_impl: &syn::ItemImpl) -> Vec<FnInfo> {
430    item_impl
431        .items
432        .iter()
433        .enumerate()
434        .filter_map(|(source_idx, item)| {
435            let syn::ImplItem::Fn(fn_item) = item else {
436                return None;
437            };
438            Some(FnInfo {
439                source_idx,
440                span: fn_item.sig.fn_token.span,
441                name: fn_item.sig.ident.to_string(),
442                visibility: VisibilityClass::from(&fn_item.vis),
443                calls: self::direct_calls(&fn_item.block, true),
444            })
445        })
446        .collect()
447}
448
449fn direct_calls(block: &syn::Block, associated: bool) -> Vec<String> {
450    let mut collector = DirectCallCollector {
451        associated,
452        calls: Vec::new(),
453    };
454    collector.visit_block(block);
455    collector.calls
456}
457
458fn direct_call_name(path: &syn::ExprPath, associated: bool) -> Option<String> {
459    // Caller order is intentionally syntax-only: resolve only local fn paths.
460    if path.qself.is_some() || path.path.leading_colon.is_some() {
461        return None;
462    }
463    let mut segments = path.path.segments.iter();
464    let first = segments.next()?;
465    if associated {
466        if first.ident == "Self"
467            && let Some(second) = segments.next()
468            && segments.next().is_none()
469        {
470            return Some(second.ident.to_string());
471        }
472        None
473    } else if segments.next().is_none() {
474        Some(first.ident.to_string())
475    } else {
476        None
477    }
478}
479
480#[cfg(test)]
481mod tests {
482    use std::path::PathBuf;
483
484    use test_that::prelude::*;
485
486    use super::MisorderedFnDetails;
487    use super::MisorderedFnRule;
488    use super::MisorderedFnViolation;
489    use crate::cmds::rsl::ast::ItemKind;
490    use crate::cmds::rsl::rules::TypedRule;
491    use crate::cmds::rsl::rules::common::Location;
492
493    #[test]
494    fn test_misordered_fn_check_when_private_helper_precedes_caller_reports_helper() {
495        let syntax = syn::parse_file(
496            r"
497            fn helper() {}
498            fn caller() {
499                helper();
500            }
501            ",
502        )
503        .unwrap();
504
505        let result = MisorderedFnRule.check(&crate::cmds::rsl::rules::test_ctx(&syntax));
506
507        assert_that!(
508            result,
509            eq(vec![MisorderedFnViolation {
510                location: Location::new(PathBuf::from("test.rs"), 2, 13),
511                details: MisorderedFnDetails {
512                    expected_after: "fn caller".to_owned(),
513                    item: ItemKind::Fn,
514                },
515            }])
516        );
517    }
518
519    #[test]
520    fn test_misordered_fn_check_when_mutually_recursive_helpers_are_split_reports_later_helper() {
521        let syntax = syn::parse_file(
522            r"
523            fn first() {
524                second();
525            }
526            fn gap() {}
527            fn second() {
528                first();
529            }
530            ",
531        )
532        .unwrap();
533
534        let result = MisorderedFnRule.check(&crate::cmds::rsl::rules::test_ctx(&syntax));
535
536        assert_that!(
537            result,
538            eq(vec![MisorderedFnViolation {
539                location: Location::new(PathBuf::from("test.rs"), 6, 13),
540                details: MisorderedFnDetails {
541                    expected_after: "fn first".to_owned(),
542                    item: ItemKind::Fn,
543                },
544            }])
545        );
546    }
547
548    #[test]
549    fn test_misordered_fn_check_when_external_module_is_declared_does_not_read_external_file() {
550        let syntax = syn::parse_file(
551            r"
552            mod external;
553            fn run() {}
554            ",
555        )
556        .unwrap();
557
558        let result = MisorderedFnRule.check(&crate::cmds::rsl::rules::test_ctx(&syntax));
559
560        assert_that!(result, is_empty());
561    }
562}