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 fn check_include_groups(
67 &self,
68 items: &[TopLevelItem],
69 file_path: &Path,
70 config_header: &str,
71 violations: &mut Vec<Violation>,
72 ) {
73 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 self.check_include_groups(body, file_path, config_header, violations);
92 }
93 _ => {}
94 }
95 }
96
97 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 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 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; }
128
129 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 let trailing_newlines = last_inc.location.count_trailing_newlines();
139
140 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 let mut fixes = Vec::new();
151
152 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 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 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 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 for _ in 0..trailing_newlines {
242 result.push('\n');
243 }
244
245 result
246 }
247
248 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, Associated = 1, StandardC = 2, System = 3, Project = 4, }
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 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 fn is_standard_c_header(&self, path: &str) -> bool {
290 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 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 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), format!("{}-private.h", stem), format!("{}private.h", stem), ]
411 }
412
413 fn is_associated_header(&self, include_path: &str, file_path: &Path) -> bool {
417 let basenames = self.get_associated_header_basenames(file_path);
418
419 let include_basename = include_path.rsplit('/').next().unwrap_or(include_path);
421
422 basenames.iter().any(|pattern| pattern == include_basename)
423 }
424}