Skip to main content

gobject_linter/rules/
include_order.rs

1use std::{path::Path, sync::LazyLock};
2
3use gobject_ast::model::{Include, PreprocessorDirective, TopLevelItem};
4
5use crate::{
6    ast_context::AstContext,
7    config::Config,
8    rules::{ConfigOption, Fix, Rule, Violation},
9};
10
11pub struct IncludeOrder;
12
13impl Rule for IncludeOrder {
14    fn name(&self) -> &'static str {
15        "include_order"
16    }
17
18    fn description(&self) -> &'static str {
19        "Enforce consistent include ordering: config header (configurable), associated header, standard C/POSIX headers, system headers, project headers"
20    }
21
22    fn category(&self) -> crate::rules::Category {
23        crate::rules::Category::Style
24    }
25
26    fn fixable(&self) -> bool {
27        true
28    }
29
30    fn config_options(&self) -> &'static [ConfigOption] {
31        static OPTIONS: LazyLock<Vec<ConfigOption>> = LazyLock::new(|| {
32            vec![ConfigOption {
33                name: "config_header",
34                option_type: "string",
35                default_value: "\"config.h\"",
36                example_value: "\"myproject-config.h\"",
37                description: "Name of the config header file",
38            }]
39        });
40
41        &OPTIONS
42    }
43
44    fn check_all(
45        &self,
46        ast_context: &AstContext,
47        config: &Config,
48        violations: &mut Vec<Violation>,
49    ) {
50        let config_header = config
51            .get_rule_config(self.name())
52            .and_then(|rc| rc.options.get("config_header"))
53            .and_then(|v| v.as_str())
54            .unwrap_or("config.h");
55
56        for (path, file) in ast_context.iter_all_files() {
57            self.check_include_groups(&file.top_level_items, path, config_header, violations);
58        }
59    }
60}
61
62impl IncludeOrder {
63    /// Check include ordering at each level of the tree structure
64    /// All includes at the same level (outside conditionals) are sorted
65    /// together
66    fn check_include_groups(
67        &self,
68        items: &[TopLevelItem],
69        file_path: &Path,
70        config_header: &str,
71        violations: &mut Vec<Violation>,
72    ) {
73        // Collect all top-level includes
74        let mut top_level_includes = Vec::new();
75
76        for item in items {
77            match item {
78                TopLevelItem::Preprocessor(PreprocessorDirective::Include {
79                    path,
80                    is_system,
81                    location,
82                }) => {
83                    top_level_includes.push(Include {
84                        path: path.clone(),
85                        is_system: *is_system,
86                        location: location.clone(),
87                    });
88                }
89                TopLevelItem::Preprocessor(PreprocessorDirective::Conditional { body, .. }) => {
90                    // Recursively check includes within the conditional block
91                    self.check_include_groups(body, file_path, config_header, violations);
92                }
93                _ => {}
94            }
95        }
96
97        // Check and fix all top-level includes as one group
98        if !top_level_includes.is_empty() {
99            self.check_and_fix_group_scattered(
100                &top_level_includes,
101                file_path,
102                config_header,
103                violations,
104            );
105        }
106    }
107
108    /// Check and fix scattered includes (may be separated by #ifdef blocks)
109    /// All includes should be sorted and moved to be consecutive at the start
110    fn check_and_fix_group_scattered(
111        &self,
112        includes: &[Include],
113        file_path: &Path,
114        config_header: &str,
115        violations: &mut Vec<Violation>,
116    ) {
117        if includes.is_empty() {
118            return;
119        }
120
121        // Step 1: Compute expected vs actual order
122        let expected_order = self.compute_expected_order(file_path, includes, config_header);
123        let current_order: Vec<_> = includes.iter().map(|inc| &inc.path).collect();
124
125        if expected_order == current_order {
126            return; // Already in correct order
127        }
128
129        // Step 2: Gather information about the include block
130        let first_inc = &includes[0];
131        let last_inc = includes
132            .iter()
133            .max_by_key(|inc| inc.location.end_byte)
134            .unwrap();
135
136        // Step 3: Determine spacing after the last include (to preserve it)
137        // The last include's end_byte points right after its newline
138        let trailing_newlines = last_inc.location.count_trailing_newlines();
139
140        // Step 4: Generate sorted includes with preserved trailing spacing
141        let sorted_text = self.generate_sorted_includes_text(
142            &expected_order,
143            includes,
144            file_path,
145            config_header,
146            trailing_newlines,
147        );
148
149        // Step 5: Build fixes
150        let mut fixes = Vec::new();
151
152        // Replace first include with all sorted includes (including trailing newlines)
153        // Also consume any blank lines immediately after the first include
154        let first_end = first_inc.location.end_byte + first_inc.location.count_trailing_newlines();
155        fixes.push(Fix::new(
156            first_inc.location.start_byte,
157            first_end,
158            sorted_text,
159        ));
160
161        // Delete all other includes
162        // Consume any blank lines immediately after each deleted include
163        for inc in includes.iter().skip(1) {
164            let end = inc.location.end_byte + inc.location.count_trailing_newlines();
165            fixes.push(Fix::delete(inc.location.start_byte, end));
166        }
167
168        violations.push(self.violation_with_fixes(
169            file_path,
170            first_inc.location.line,
171            1,
172            format!(
173                "Includes are not in the correct order. Expected: {} (if present), associated header, standard C headers, system headers (<>), project headers (\"\") (all alphabetically sorted within each group, blank line between groups)",
174                config_header
175            ),
176            fixes,
177        ));
178    }
179
180    /// Generate the text for sorted includes with proper grouping
181    fn generate_sorted_includes_text(
182        &self,
183        expected_order: &[&str],
184        includes: &[Include],
185        file_path: &Path,
186        config_header: &str,
187        trailing_newlines: usize,
188    ) -> String {
189        #[derive(Debug, Clone, Copy, PartialEq, Eq)]
190        enum IncludeGroup {
191            Config,
192            Associated,
193            StandardC,
194            System,
195            Project,
196        }
197
198        let grouped_includes: Vec<(&&str, IncludeGroup)> = expected_order
199            .iter()
200            .map(|path| {
201                let group = if *path == config_header {
202                    IncludeGroup::Config
203                } else if self.is_associated_header(path, file_path) {
204                    IncludeGroup::Associated
205                } else {
206                    let original = includes.iter().find(|inc| inc.path == *path).unwrap();
207                    if original.is_system && self.is_standard_c_header(path) {
208                        IncludeGroup::StandardC
209                    } else if original.is_system {
210                        IncludeGroup::System
211                    } else {
212                        IncludeGroup::Project
213                    }
214                };
215                (path, group)
216            })
217            .collect();
218
219        let mut result = String::new();
220        let mut last_group: Option<IncludeGroup> = None;
221
222        for (path, group) in grouped_includes {
223            // Add blank line between groups
224            if let Some(prev_group) = last_group
225                && prev_group != group
226            {
227                result.push('\n');
228            }
229            last_group = Some(group);
230
231            let original = includes.iter().find(|inc| inc.path == *path).unwrap();
232            let bracket = if original.is_system {
233                ("<", ">")
234            } else {
235                ("\"", "\"")
236            };
237            result.push_str(&format!("#include {}{}{}\n", bracket.0, path, bracket.1));
238        }
239
240        // Add trailing newlines to preserve original spacing
241        for _ in 0..trailing_newlines {
242            result.push('\n');
243        }
244
245        result
246    }
247
248    /// Compute the expected order of includes
249    fn compute_expected_order<'a>(
250        &self,
251        file_path: &Path,
252        includes: &'a [Include],
253        config_header: &str,
254    ) -> Vec<&'a str> {
255        #[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
256        enum IncludeGroup {
257            Config = 0,     // config.h must be first
258            Associated = 1, // foo.c -> foo.h
259            StandardC = 2,  // <stdio.h>, <math.h>, etc.
260            System = 3,     // <glib.h>, <gtk/gtk.h>, etc.
261            Project = 4,    // "..."
262        }
263
264        let mut grouped: Vec<(&Include, IncludeGroup)> = includes
265            .iter()
266            .map(|inc| {
267                let group = if inc.path == config_header {
268                    IncludeGroup::Config
269                } else if self.is_associated_header(&inc.path, file_path) {
270                    IncludeGroup::Associated
271                } else if inc.is_system && self.is_standard_c_header(&inc.path) {
272                    IncludeGroup::StandardC
273                } else if inc.is_system {
274                    IncludeGroup::System
275                } else {
276                    IncludeGroup::Project
277                };
278                (inc, group)
279            })
280            .collect();
281
282        // Sort by group first, then alphabetically within each group
283        grouped.sort_by(|a, b| a.1.cmp(&b.1).then_with(|| a.0.path.cmp(&b.0.path)));
284
285        grouped.iter().map(|(inc, _)| inc.path.as_str()).collect()
286    }
287
288    /// Check if a header is a standard C or POSIX header
289    fn is_standard_c_header(&self, path: &str) -> bool {
290        // Standard C library headers
291        if matches!(
292            path,
293            "assert.h"
294                | "complex.h"
295                | "ctype.h"
296                | "errno.h"
297                | "fenv.h"
298                | "float.h"
299                | "inttypes.h"
300                | "iso646.h"
301                | "limits.h"
302                | "locale.h"
303                | "math.h"
304                | "setjmp.h"
305                | "signal.h"
306                | "stdalign.h"
307                | "stdarg.h"
308                | "stdatomic.h"
309                | "stdbool.h"
310                | "stddef.h"
311                | "stdint.h"
312                | "stdio.h"
313                | "stdlib.h"
314                | "stdnoreturn.h"
315                | "string.h"
316                | "tgmath.h"
317                | "threads.h"
318                | "time.h"
319                | "uchar.h"
320                | "wchar.h"
321                | "wctype.h"
322        ) {
323            return true;
324        }
325
326        // POSIX headers
327        if matches!(
328            path,
329            "aio.h"
330                | "arpa/inet.h"
331                | "dirent.h"
332                | "dlfcn.h"
333                | "fcntl.h"
334                | "fmtmsg.h"
335                | "fnmatch.h"
336                | "ftw.h"
337                | "glob.h"
338                | "grp.h"
339                | "iconv.h"
340                | "langinfo.h"
341                | "libgen.h"
342                | "monetary.h"
343                | "mqueue.h"
344                | "ndbm.h"
345                | "net/if.h"
346                | "netdb.h"
347                | "netinet/in.h"
348                | "netinet/tcp.h"
349                | "nl_types.h"
350                | "poll.h"
351                | "pthread.h"
352                | "pwd.h"
353                | "regex.h"
354                | "sched.h"
355                | "search.h"
356                | "semaphore.h"
357                | "spawn.h"
358                | "strings.h"
359                | "stropts.h"
360                | "sys/ipc.h"
361                | "sys/mman.h"
362                | "sys/msg.h"
363                | "sys/resource.h"
364                | "sys/select.h"
365                | "sys/sem.h"
366                | "sys/shm.h"
367                | "sys/socket.h"
368                | "sys/stat.h"
369                | "sys/statvfs.h"
370                | "sys/time.h"
371                | "sys/times.h"
372                | "sys/types.h"
373                | "sys/uio.h"
374                | "sys/un.h"
375                | "sys/utsname.h"
376                | "sys/wait.h"
377                | "sysexits.h"
378                | "syslog.h"
379                | "tar.h"
380                | "termios.h"
381                | "trace.h"
382                | "ulimit.h"
383                | "unistd.h"
384                | "utime.h"
385                | "utmpx.h"
386                | "wordexp.h"
387        ) {
388            return true;
389        }
390
391        false
392    }
393
394    /// Get all possible associated headers for a C file
395    /// Returns the basenames to check (without directory prefix)
396    /// foo.c -> ["foo.h", "foo-private.h", "fooprivate.h"]
397    fn get_associated_header_basenames(&self, file_path: &Path) -> Vec<String> {
398        if file_path.extension() != Some(std::ffi::OsStr::new("c")) {
399            return Vec::new();
400        }
401
402        let Some(stem) = file_path.file_stem().and_then(|s| s.to_str()) else {
403            return Vec::new();
404        };
405
406        vec![
407            format!("{}.h", stem),         // foo.c -> foo.h
408            format!("{}-private.h", stem), // foo.c -> foo-private.h
409            format!("{}private.h", stem),  // foo.c -> fooprivate.h
410        ]
411    }
412
413    /// Check if an include path is an associated header for the given file
414    /// Checks the basename of the include, so "meta-test/meta-test-monitor.h"
415    /// matches for "meta-test-monitor.c"
416    fn is_associated_header(&self, include_path: &str, file_path: &Path) -> bool {
417        let basenames = self.get_associated_header_basenames(file_path);
418
419        // Extract basename from include path (part after last '/')
420        let include_basename = include_path.rsplit('/').next().unwrap_or(include_path);
421
422        basenames.iter().any(|pattern| pattern == include_basename)
423    }
424}