diff --git a/python/src/com/jetbrains/python/actions/RemoveArgumentEqualDefaultQuickFix.java b/python/src/com/jetbrains/python/actions/RemoveArgumentEqualDefaultQuickFix.java index dd3c22c4fd73..ca0d8e2115f6 100644 --- a/python/src/com/jetbrains/python/actions/RemoveArgumentEqualDefaultQuickFix.java +++ b/python/src/com/jetbrains/python/actions/RemoveArgumentEqualDefaultQuickFix.java @@ -3,14 +3,15 @@ package com.jetbrains.python.actions; import com.intellij.codeInspection.LocalQuickFix; import com.intellij.codeInspection.ProblemDescriptor; import com.intellij.openapi.project.Project; +import com.intellij.openapi.util.text.StringUtil; import com.intellij.psi.PsiElement; import com.intellij.psi.util.PsiTreeUtil; import com.jetbrains.python.PyBundle; -import com.jetbrains.python.psi.PyArgumentList; -import com.jetbrains.python.psi.PyExpression; -import com.jetbrains.python.psi.PyKeywordArgument; +import com.jetbrains.python.psi.*; import org.jetbrains.annotations.NotNull; +import java.util.ArrayList; +import java.util.List; import java.util.Set; /** @@ -36,14 +37,23 @@ public class RemoveArgumentEqualDefaultQuickFix implements LocalQuickFix { public void applyFix(@NotNull Project project, @NotNull ProblemDescriptor descriptor) { PsiElement element = descriptor.getPsiElement(); - PyExpression[] arguments = PsiTreeUtil.getParentOfType(element, PyArgumentList.class).getArguments(); - boolean canDelete = true; - for (int i = arguments.length-1; i != -1; --i) { - if (myProblemElements.contains(arguments[i])) { - if (canDelete) - arguments[i].delete(); + + PyArgumentList argumentList = PsiTreeUtil.getParentOfType(element, PyArgumentList.class); + if (argumentList == null) return; + StringBuilder newArgumentList = new StringBuilder("foo("); + + PyExpression[] arguments = argumentList.getArguments(); + List newArgs = new ArrayList(); + for (int i = 0; i != arguments.length; ++i) { + if (!myProblemElements.contains(arguments[i])) { + newArgs.add(arguments[i].getText()); } - else if (!(arguments[i] instanceof PyKeywordArgument)) canDelete = false; } + + newArgumentList.append(StringUtil.join(newArgs, ", ")).append(")"); + PyExpression expression = PyElementGenerator.getInstance(project).createFromText( + LanguageLevel.forElement(argumentList), PyExpressionStatement.class, newArgumentList.toString()).getExpression(); + if (expression instanceof PyCallExpression) + argumentList.replace(((PyCallExpression)expression).getArgumentList()); } } diff --git a/python/src/com/jetbrains/python/inspections/PyRedundantParenthesesInspection.java b/python/src/com/jetbrains/python/inspections/PyRedundantParenthesesInspection.java index e81287e52f25..2321395c7445 100644 --- a/python/src/com/jetbrains/python/inspections/PyRedundantParenthesesInspection.java +++ b/python/src/com/jetbrains/python/inspections/PyRedundantParenthesesInspection.java @@ -69,6 +69,8 @@ public class PyRedundantParenthesesInspection extends PyInspection { registerProblem(node, "Remove redundant parentheses", new RedundantParenthesesQuickFix()); } else if (expression instanceof PyBinaryExpression) { + if (node.getParent() instanceof PyPrefixExpression) + return; if (((PyBinaryExpression)expression).getOperator() == PyTokenTypes.AND_KEYWORD || ((PyBinaryExpression)expression).getOperator() == PyTokenTypes.OR_KEYWORD) { if (((PyBinaryExpression)expression).getLeftExpression() instanceof PyParenthesizedExpression && diff --git a/python/testData/inspections/ArgumentEqualDefault.py b/python/testData/inspections/ArgumentEqualDefault.py index 275e16df0bd2..2db19b255ef1 100644 --- a/python/testData/inspections/ArgumentEqualDefault.py +++ b/python/testData/inspections/ArgumentEqualDefault.py @@ -2,4 +2,5 @@ def foo(a, b = 345, c = 1): pass #PY-3261 -foo(1, 345, c=22) \ No newline at end of file +foo(1, +345, c=22) \ No newline at end of file diff --git a/python/testData/inspections/PyArgumentEqualDefaultInspection/test.py b/python/testData/inspections/PyArgumentEqualDefaultInspection/test.py index f7f0da3a04e6..f95271a0abba 100644 --- a/python/testData/inspections/PyArgumentEqualDefaultInspection/test.py +++ b/python/testData/inspections/PyArgumentEqualDefaultInspection/test.py @@ -58,3 +58,12 @@ def bar(a = "qwer"): bar(a = 'qwer') getattr(bar, "__doc__", None) # None is not highlighted + +class a: + def get(self, a, b = None): + pass + +kw = a() + +kw['customerPaymentProfileId'] = kw.get("customerPaymentProfileId", + None)