Skip to main content

gobject_linter/rules/
g_object_virtual_methods_chain_up.rs

1use gobject_ast::model::{Expression, Parameter, Statement};
2
3use crate::{
4    ast_context::AstContext,
5    config::Config,
6    rules::{Category, Fix, Rule, Violation},
7};
8
9const CHAINABLE_VFUNCS: &[&str] = &["dispose", "finalize", "constructed"];
10
11pub struct GObjectVirtualMethodsChainUp;
12
13impl Rule for GObjectVirtualMethodsChainUp {
14    fn name(&self) -> &'static str {
15        "g_object_virtual_methods_chain_up"
16    }
17
18    fn description(&self) -> &'static str {
19        "Ensure dispose/finalize/constructed methods chain up to parent class"
20    }
21
22    fn category(&self) -> Category {
23        Category::Correctness
24    }
25
26    fn fixable(&self) -> bool {
27        true
28    }
29
30    fn check_all(
31        &self,
32        ast_context: &AstContext,
33        config: &Config,
34        violations: &mut Vec<Violation>,
35    ) {
36        for (_path, file) in ast_context.iter_all_files() {
37            for gt in file
38                .iter_all_gobject_types()
39                .filter(|gt| gt.kind.is_define())
40            {
41                let vfuncs = file.resolve_class_init_vfuncs(gt);
42
43                for ((class_type, field), func_name) in &vfuncs {
44                    if class_type != "GObjectClass" || !CHAINABLE_VFUNCS.contains(field) {
45                        continue;
46                    }
47
48                    let Some(func) = file
49                        .iter_function_definitions()
50                        .find(|f| f.name == *func_name)
51                    else {
52                        continue;
53                    };
54
55                    if has_chainup_call(&func.body_statements, field) {
56                        continue;
57                    }
58
59                    let param_name = func
60                        .parameters
61                        .first()
62                        .and_then(|p| match p {
63                            Parameter::Regular { name, .. } => name.as_deref(),
64                            _ => None,
65                        })
66                        .unwrap_or("object");
67
68                    let parent_class = format!("{}_parent_class", gt.function_prefix);
69                    let cast = config.style.format_call("G_OBJECT_CLASS", &[&parent_class]);
70                    let args = config.style.format_call("", &[param_name]);
71                    let chainup_call = format!("{cast}->{field}{args};");
72
73                    let fix = if let Some(body_loc) = &func.body_location {
74                        let indent = func
75                            .body_statements
76                            .iter()
77                            .find(|s| {
78                                matches!(s, Statement::Declaration(_) | Statement::Expression(_))
79                            })
80                            .map_or_else(
81                                || "  ".to_string(),
82                                |s| s.location().extract_line_indentation(),
83                            );
84
85                        let (pos, before, after) = if *field == "constructed" {
86                            let last_decl = func
87                                .body_statements
88                                .iter()
89                                .filter_map(|s| match s {
90                                    Statement::Declaration(d) => Some(&d.location),
91                                    _ => None,
92                                })
93                                .next_back();
94
95                            if let Some(loc) = last_decl {
96                                let after = if loc.count_trailing_newlines() >= 2 {
97                                    ""
98                                } else {
99                                    "\n"
100                                };
101                                (loc.end_byte, "\n\n", after)
102                            } else {
103                                (body_loc.start_byte + 1, "\n", "\n")
104                            }
105                        } else {
106                            let trailing = func
107                                .body_statements
108                                .last()
109                                .map_or(1, |s| s.location().count_trailing_newlines());
110                            let before = if trailing >= 2 { "" } else { "\n" };
111                            (body_loc.end_byte - 1, before, "\n")
112                        };
113
114                        Some(Fix::new(
115                            pos,
116                            pos,
117                            format!("{before}{indent}{chainup_call}{after}"),
118                        ))
119                    } else {
120                        None
121                    };
122
123                    let msg = format!(
124                        "{func_name} must chain up to parent class (e.g., {chainup_call})",
125                    );
126                    let violation = if let Some(fix) = fix {
127                        self.violation_with_fix(
128                            &file.path,
129                            func.location.line,
130                            func.location.column,
131                            msg,
132                            fix,
133                        )
134                    } else {
135                        self.violation(&file.path, func.location.line, func.location.column, msg)
136                    };
137                    violations.push(violation);
138                }
139            }
140        }
141    }
142}
143
144fn has_chainup_call(statements: &[Statement], method_type: &str) -> bool {
145    for stmt in statements {
146        let mut found = false;
147        stmt.walk(&mut |s| {
148            s.visit_expressions(&mut |expr| {
149                expr.walk(&mut |e| {
150                    if let Expression::Call(call) = e {
151                        let func = match &*call.function {
152                            Expression::Unary(u) => &u.operand,
153                            other => other,
154                        };
155                        if let Expression::FieldAccess(fa) = func
156                            && fa.field == method_type
157                        {
158                            found = true;
159                        }
160                    }
161                });
162            });
163        });
164        if found {
165            return true;
166        }
167    }
168    false
169}