diff --git a/crates/ruff_linter/src/rules/flake8_simplify/rules/needless_bool.rs b/crates/ruff_linter/src/rules/flake8_simplify/rules/needless_bool.rs index 4a7d6bedc6..fc64997661 100644 --- a/crates/ruff_linter/src/rules/flake8_simplify/rules/needless_bool.rs +++ b/crates/ruff_linter/src/rules/flake8_simplify/rules/needless_bool.rs @@ -5,6 +5,7 @@ use ruff_python_semantic::analyze::typing::{is_sys_version_block, is_type_checki use ruff_text_size::{Ranged, TextRange}; use crate::checkers::ast::Checker; +use crate::fix::snippet::SourceCodeSnippet; /// ## What it does /// Checks for `if` statements that can be replaced with `bool`. @@ -30,7 +31,8 @@ use crate::checkers::ast::Checker; /// - [Python documentation: Truth Value Testing](https://docs.python.org/3/library/stdtypes.html#truth-value-testing) #[violation] pub struct NeedlessBool { - condition: String, + condition: SourceCodeSnippet, + replacement: Option, } impl Violation for NeedlessBool { @@ -38,13 +40,24 @@ impl Violation for NeedlessBool { #[derive_message_formats] fn message(&self) -> String { - let NeedlessBool { condition } = self; - format!("Return the condition `{condition}` directly") + let NeedlessBool { condition, .. } = self; + if let Some(condition) = condition.full_display() { + format!("Return the condition `{condition}` directly") + } else { + format!("Return the condition directly") + } } fn fix_title(&self) -> Option { - let NeedlessBool { condition } = self; - Some(format!("Replace with `return {condition}`")) + let NeedlessBool { replacement, .. } = self; + if let Some(replacement) = replacement + .as_ref() + .and_then(SourceCodeSnippet::full_display) + { + Some(format!("Replace with `{replacement}`")) + } else { + Some(format!("Inline condition")) + } } } @@ -90,11 +103,20 @@ pub(crate) fn needless_bool(checker: &mut Checker, stmt_if: &ast::StmtIf) { return; }; - // If the branches have the same condition, abort (although the code could be - // simplified). - if if_return == else_return { - return; - } + // Determine whether the return values are inverted, as in: + // ```python + // if x > 0: + // return False + // else: + // return True + // ``` + let inverted = match (if_return, else_return) { + (Bool::True, Bool::False) => false, + (Bool::False, Bool::True) => true, + // If the branches have the same condition, abort (although the code could be + // simplified). + _ => return, + }; // Avoid suggesting ternary for `if sys.version_info >= ...`-style checks. if is_sys_version_block(stmt_if, checker.semantic()) { @@ -106,33 +128,38 @@ pub(crate) fn needless_bool(checker: &mut Checker, stmt_if: &ast::StmtIf) { return; } - let condition = checker.generator().expr(if_test); - let mut diagnostic = Diagnostic::new(NeedlessBool { condition }, range); - if matches!(if_return, Bool::True) - && matches!(else_return, Bool::False) - && !checker.indexer().has_comments(&range, checker.locator()) - && (if_test.is_compare_expr() || checker.semantic().is_builtin("bool")) - { - if if_test.is_compare_expr() { - // If the condition is a comparison, we can replace it with the condition. + let condition = checker.locator().slice(if_test); + let replacement = if checker.indexer().has_comments(&range, checker.locator()) { + None + } else { + // If the return values are inverted, wrap the condition in a `not`. + if inverted { + let node = ast::StmtReturn { + value: Some(Box::new(Expr::UnaryOp(ast::ExprUnaryOp { + op: ast::UnaryOp::Not, + operand: Box::new(if_test.clone()), + range: TextRange::default(), + }))), + range: TextRange::default(), + }; + Some(checker.generator().stmt(&node.into())) + } else if if_test.is_compare_expr() { + // If the condition is a comparison, we can replace it with the condition, since we + // know it's a boolean. let node = ast::StmtReturn { value: Some(Box::new(if_test.clone())), range: TextRange::default(), }; - diagnostic.set_fix(Fix::unsafe_edit(Edit::range_replacement( - checker.generator().stmt(&node.into()), - range, - ))); - } else { - // Otherwise, we need to wrap the condition in a call to `bool`. (We've already - // verified, above, that `bool` is a builtin.) - let node = ast::ExprName { + Some(checker.generator().stmt(&node.into())) + } else if checker.semantic().is_builtin("bool") { + // Otherwise, we need to wrap the condition in a call to `bool`. + let func_node = ast::ExprName { id: "bool".into(), ctx: ExprContext::Load, range: TextRange::default(), }; - let node1 = ast::ExprCall { - func: Box::new(node.into()), + let value_node = ast::ExprCall { + func: Box::new(func_node.into()), arguments: Arguments { args: vec![if_test.clone()], keywords: vec![], @@ -140,15 +167,28 @@ pub(crate) fn needless_bool(checker: &mut Checker, stmt_if: &ast::StmtIf) { }, range: TextRange::default(), }; - let node2 = ast::StmtReturn { - value: Some(Box::new(node1.into())), + let return_node = ast::StmtReturn { + value: Some(Box::new(value_node.into())), range: TextRange::default(), }; - diagnostic.set_fix(Fix::unsafe_edit(Edit::range_replacement( - checker.generator().stmt(&node2.into()), - range, - ))); - }; + Some(checker.generator().stmt(&return_node.into())) + } else { + None + } + }; + + let mut diagnostic = Diagnostic::new( + NeedlessBool { + condition: SourceCodeSnippet::from_str(condition), + replacement: replacement.clone().map(SourceCodeSnippet::new), + }, + range, + ); + if let Some(replacement) = replacement { + diagnostic.set_fix(Fix::unsafe_edit(Edit::range_replacement( + replacement, + range, + ))); } checker.diagnostics.push(diagnostic); } diff --git a/crates/ruff_linter/src/rules/flake8_simplify/snapshots/ruff_linter__rules__flake8_simplify__tests__SIM103_SIM103.py.snap b/crates/ruff_linter/src/rules/flake8_simplify/snapshots/ruff_linter__rules__flake8_simplify__tests__SIM103_SIM103.py.snap index 929549a3e3..174d22184c 100644 --- a/crates/ruff_linter/src/rules/flake8_simplify/snapshots/ruff_linter__rules__flake8_simplify__tests__SIM103_SIM103.py.snap +++ b/crates/ruff_linter/src/rules/flake8_simplify/snapshots/ruff_linter__rules__flake8_simplify__tests__SIM103_SIM103.py.snap @@ -12,7 +12,7 @@ SIM103.py:3:5: SIM103 [*] Return the condition `a` directly 6 | | return False | |____________________^ SIM103 | - = help: Replace with `return a` + = help: Replace with `return bool(a)` ℹ Unsafe fix 1 1 | def f(): @@ -63,7 +63,7 @@ SIM103.py:21:5: SIM103 [*] Return the condition `b` directly 24 | | return False | |____________________^ SIM103 | - = help: Replace with `return b` + = help: Replace with `return bool(b)` ℹ Unsafe fix 18 18 | # SIM103 @@ -89,7 +89,7 @@ SIM103.py:32:9: SIM103 [*] Return the condition `b` directly 35 | | return False | |________________________^ SIM103 | - = help: Replace with `return b` + = help: Replace with `return bool(b)` ℹ Unsafe fix 29 29 | if a: @@ -104,7 +104,7 @@ SIM103.py:32:9: SIM103 [*] Return the condition `b` directly 37 34 | 38 35 | def f(): -SIM103.py:57:5: SIM103 Return the condition `a` directly +SIM103.py:57:5: SIM103 [*] Return the condition `a` directly | 55 | def f(): 56 | # SIM103 (but not fixable) @@ -115,7 +115,20 @@ SIM103.py:57:5: SIM103 Return the condition `a` directly 60 | | return True | |___________________^ SIM103 | - = help: Replace with `return a` + = help: Replace with `return not a` + +ℹ Unsafe fix +54 54 | +55 55 | def f(): +56 56 | # SIM103 (but not fixable) +57 |- if a: +58 |- return False +59 |- else: +60 |- return True + 57 |+ return not a +61 58 | +62 59 | +63 60 | def f(): SIM103.py:83:5: SIM103 Return the condition `a` directly | @@ -128,6 +141,6 @@ SIM103.py:83:5: SIM103 Return the condition `a` directly 86 | | return False | |____________________^ SIM103 | - = help: Replace with `return a` + = help: Inline condition