use std::collections::HashSet;
use gobject_ast::{Expression, Statement};
use crate::{
ast_context::AstContext,
config::Config,
rules::{Fix, Rule, Violation},
};
pub struct UseGSourceConstants;
impl Rule for UseGSourceConstants {
fn name(&self) -> &'static str {
"use_g_source_constants"
}
fn description(&self) -> &'static str {
"Use G_SOURCE_CONTINUE/G_SOURCE_REMOVE instead of TRUE/FALSE in GSourceFunc callbacks"
}
fn category(&self) -> crate::rules::Category {
crate::rules::Category::Style
}
fn fixable(&self) -> bool {
true
}
fn check_all(
&self,
ast_context: &AstContext,
_config: &Config,
violations: &mut Vec<Violation>,
) {
let mut callbacks: HashSet<&str> = HashSet::new();
for (_path, file) in ast_context.iter_c_files() {
for func in file.iter_function_definitions() {
for call in func.find_calls(&[
"g_idle_add",
"g_idle_add_full",
"g_timeout_add",
"g_timeout_add_seconds",
"g_timeout_add_full",
"g_timeout_add_seconds_full",
]) {
if let Some(name) = self.extract_callback_name(call, &file.source) {
callbacks.insert(name);
}
}
}
}
if callbacks.is_empty() {
return;
}
for (path, file) in ast_context.iter_c_files() {
for func in file.iter_function_definitions() {
if callbacks.contains(func.name.as_str()) {
self.check_statements(path, &func.body_statements, &file.source, violations);
}
}
}
}
}
impl UseGSourceConstants {
fn extract_callback_name<'a>(
&self,
call: &'a gobject_ast::CallExpression,
_source: &[u8],
) -> Option<&'a str> {
let func_name = call.function_name_str()?;
let callback_arg_index: usize = match func_name {
"g_idle_add" => 0,
"g_idle_add_full" | "g_timeout_add" | "g_timeout_add_seconds" => 1,
"g_timeout_add_full" | "g_timeout_add_seconds_full" => 2,
_ => return None,
};
if callback_arg_index >= call.arguments.len() {
return None;
}
let arg_expr = call.get_arg(callback_arg_index)?;
if let Expression::Identifier(id) = arg_expr {
Some(&id.name)
} else {
None
}
}
fn check_statements(
&self,
file_path: &std::path::Path,
statements: &[Statement],
source: &[u8],
violations: &mut Vec<Violation>,
) {
for stmt in statements {
for ret_stmt in stmt.iter_returns() {
if let Some(value) = &ret_stmt.value {
self.check_return_value(file_path, value, source, violations);
}
}
}
}
fn check_return_value(
&self,
file_path: &std::path::Path,
expr: &Expression,
_source: &[u8],
violations: &mut Vec<Violation>,
) {
expr.walk(&mut |e| match e {
Expression::Identifier(id) if id.name == "TRUE" || id.name == "FALSE" => {
let replacement = if id.name == "TRUE" {
"G_SOURCE_CONTINUE"
} else {
"G_SOURCE_REMOVE"
};
let fix = Fix::new(
id.location.start_byte,
id.location.end_byte,
replacement.to_string(),
);
violations.push(self.violation_with_fix(
file_path,
id.location.line,
id.location.column,
format!(
"Use {} instead of {} in GSourceFunc callback",
replacement, id.name
),
fix,
));
}
Expression::Boolean(b) => {
let (old_name, replacement) = if b.value {
("TRUE", "G_SOURCE_CONTINUE")
} else {
("FALSE", "G_SOURCE_REMOVE")
};
let fix = Fix::new(
b.location.start_byte,
b.location.end_byte,
replacement.to_string(),
);
violations.push(self.violation_with_fix(
file_path,
b.location.line,
b.location.column,
format!(
"Use {} instead of {} in GSourceFunc callback",
replacement, old_name
),
fix,
));
}
_ => {}
});
}
}