diff --git a/python/src/com/jetbrains/python/inspections/quickfix/AddCallSuperQuickFix.java b/python/src/com/jetbrains/python/inspections/quickfix/AddCallSuperQuickFix.java index 0243cf6c0463..52087e9fbcb3 100644 --- a/python/src/com/jetbrains/python/inspections/quickfix/AddCallSuperQuickFix.java +++ b/python/src/com/jetbrains/python/inspections/quickfix/AddCallSuperQuickFix.java @@ -18,7 +18,11 @@ package com.jetbrains.python.inspections.quickfix; import com.intellij.codeInspection.LocalQuickFix; import com.intellij.codeInspection.ProblemDescriptor; import com.intellij.openapi.project.Project; +import com.intellij.openapi.util.Condition; +import com.intellij.openapi.util.Couple; +import com.intellij.openapi.util.text.StringUtil; import com.intellij.psi.util.PsiTreeUtil; +import com.intellij.util.containers.ContainerUtil; import com.jetbrains.python.PyBundle; import com.jetbrains.python.PyNames; import com.jetbrains.python.psi.*; @@ -63,7 +67,7 @@ public class AddCallSuperQuickFix implements LocalQuickFix { final PyClass superClass = superClasses[0]; final PyFunction superInit = superClass.findMethodByName(PyNames.INIT, true); if (superInit == null) return; - boolean addComma = true; + final boolean addComma; if (klass.isNewStyleClass()) { addComma = false; if (LanguageLevel.forElement(klass).isPy3K()) { @@ -74,14 +78,28 @@ public class AddCallSuperQuickFix implements LocalQuickFix { } } else { + addComma = true; superCall.append(superClass.getName()); superCall.append(".__init__(self"); } final StringBuilder newFunction = new StringBuilder("def __init__(self"); - buildParameterList(problemFunction, superInit, superCall, newFunction, addComma); + final Couple> couple = buildNewFunctionParamsAndSuperInitCallArgs(problemFunction, superInit); + final List newParameters = couple.getFirst(); + if (!newParameters.isEmpty()) { + newFunction.append(", "); + } + StringUtil.join(newParameters, ", ", newFunction); + newFunction.append("):\n\t"); + + final List superCallArguments = couple.getSecond(); + if (addComma && !superCallArguments.isEmpty()) { + superCall.append(", "); + } + StringUtil.join(superCallArguments, ", ", superCall); superCall.append(")"); + final PyStatementList statementList = problemFunction.getStatementList(); PyExpression docstring = null; final PyStatement[] statements = statementList.getStatements(); @@ -92,7 +110,6 @@ public class AddCallSuperQuickFix implements LocalQuickFix { } } - newFunction.append("):\n\t"); if (docstring != null) { newFunction.append(docstring.getText()).append("\n\t"); } @@ -110,94 +127,216 @@ public class AddCallSuperQuickFix implements LocalQuickFix { problemFunction.replace(generator.createFromText(LanguageLevel.forElement(problemFunction), PyFunction.class, newFunction.toString())); } - private static void buildParameterList(@NotNull PyFunction problemFunction, - @NotNull PyFunction superInit, - @NotNull StringBuilder superCall, - @NotNull StringBuilder newFunction, - boolean addComma) { - final PyParameter[] parameters = problemFunction.getParameterList().getParameters(); - final List problemParams = new ArrayList(); - final List functionParams = new ArrayList(); - String starName = null; - String doubleStarName = null; - for (int i = 1; i != parameters.length; i++) { - final PyParameter p = parameters[i]; - functionParams.add(p.getName()); - if (p.getText().startsWith("**")) { - doubleStarName = p.getText(); - continue; + @NotNull + private static Couple> buildNewFunctionParamsAndSuperInitCallArgs(@NotNull PyFunction origInit, + @NotNull PyFunction superInit) { + final List newFunctionParams = new ArrayList(); + final List superCallArgs = new ArrayList(); + + final ParametersInfo origInfo = new ParametersInfo(origInit.getParameterList()); + final ParametersInfo superInfo = new ParametersInfo(superInit.getParameterList()); + + // Required parameters (not-keyword) + for (PyParameter param : origInfo.getRequiredParameters()) { + newFunctionParams.add(param.getText()); + } + for (PyParameter param : superInfo.getRequiredParameters()) { + if (!origInfo.containsRequiredParam(param.getName())) { + newFunctionParams.add(param.getText()); } - if (p.getText().startsWith("*")) { - starName = p.getText(); - continue; - } - if (p.getDefaultValue() != null) { - problemParams.add(p.getText()); - continue; - } - newFunction.append(",").append(p.getText()); + superCallArgs.add(param.getName()); } - addParametersFromSuper(superInit, superCall, newFunction, problemParams, functionParams, starName, doubleStarName, addComma); + // Optional parameters (not-keyword) + for (PyParameter param : origInfo.getOptionalParameters()) { + newFunctionParams.add(param.getText()); + } + + // Positional vararg + PyParameter starredParam = null; + if (origInfo.getPositionalContainerParameter() != null) { + starredParam = origInfo.getPositionalContainerParameter(); + } + else if (superInfo.getPositionalContainerParameter() != null) { + starredParam = superInfo.getPositionalContainerParameter(); + } + else if (origInfo.getSingleStarParameter() != null) { + starredParam = origInfo.getSingleStarParameter(); + } + else if (superInfo.getSingleStarParameter() != null) { + starredParam = superInfo.getSingleStarParameter(); + } + if (starredParam != null) { + newFunctionParams.add(starredParam.getText()); + if (superInfo.getPositionalContainerParameter() != null) { + superCallArgs.add("*" + starredParam.getName()); + } + } + + // Required keyword-only parameters + for (PyParameter param : origInfo.getRequiredKeywordOnlyParameters()) { + newFunctionParams.add(param.getText()); + } + for (PyParameter param : superInfo.getRequiredKeywordOnlyParameters()) { + if (!origInfo.containsRequiredKeywordOnlyParameter(param.getName())) { + newFunctionParams.add(param.getText()); + } + superCallArgs.add(param.getName() + "=" + param.getName()); + } + + // Optional keyword-only parameters + for (PyParameter param : origInfo.getOptionalKeywordOnlyParameters()) { + newFunctionParams.add(param.getText()); + } + + // Keyword vararg + PyParameter doubleStarredParam = null; + if (origInfo.getKeywordContainerParameter() != null) { + doubleStarredParam = origInfo.getKeywordContainerParameter(); + } + else if (superInfo.getKeywordContainerParameter() != null) { + doubleStarredParam = superInfo.getKeywordContainerParameter(); + } + if (doubleStarredParam != null) { + newFunctionParams.add(doubleStarredParam.getText()); + if (superInfo.getKeywordContainerParameter() != null) { + superCallArgs.add("**" + doubleStarredParam.getName()); + } + } + return Couple.of(newFunctionParams, superCallArgs); } - private static void addParametersFromSuper(@NotNull PyFunction superInit, - @NotNull StringBuilder superCall, - @NotNull StringBuilder newFunction, - @NotNull List problemParams, - @NotNull List functionParams, - @Nullable String starName, - @Nullable String doubleStarName, - boolean addComma) { - final PyParameterList paramList = superInit.getParameterList(); - final PyParameter[] parameters = paramList.getParameters(); - boolean addDouble = false; - boolean addStar = false; - for (int i = 1; i != parameters.length; i++) { - final PyParameter p = parameters[i]; - if (p.getDefaultValue() != null) continue; - final String param = p.getName(); - final String paramText = p.getText(); - if (paramText.startsWith("**")) { - addDouble = true; - if (doubleStarName == null) { - doubleStarName = p.getText(); + private static class ParametersInfo { + + private final PyParameter mySelfParam; + /** + * Parameters without default value that come before first "*..." parameter. + */ + private final List myRequiredParams = new ArrayList(); + /** + * Parameters with default value that come before first "*..." parameter. + */ + private final List myOptionalParams = new ArrayList(); + /** + * Parameter of form "*args" (positional vararg), not the same as single "*". + */ + private final PyParameter myPositionalContainerParam; + /** + * Parameter "*", that is used to delimit normal and keyword-only parameters. + */ + private final PyParameter mySingleStarParam; + /** + * Parameters without default value that come after first "*..." parameter. + */ + private final List myRequiredKwOnlyParams = new ArrayList(); + /** + * Parameters with default value that come after first "*..." parameter. + */ + private final List myOptionalKwOnlyParams = new ArrayList(); + /** + * Parameter of form "**kwargs" (keyword vararg). + */ + private final PyParameter myKeywordContainerParam; + + public ParametersInfo(@NotNull PyParameterList parameterList) { + PyParameter positionalContainer = null; + PyParameter singleStarParam = null; + PyParameter keywordContainer = null; + PyParameter selfParam = null; + + for (PyParameter param : parameterList.getParameters()) { + if (param.isSelf()) { + selfParam = param; } - continue; - } - if (paramText.startsWith("*")) { - addStar = true; - if (starName == null) { - starName = p.getText(); + else if (param.getText().equals("*")) { + singleStarParam = param; + } + else if (param.getText().startsWith("**")) { + keywordContainer = param; + } + else if (param.getText().startsWith("*")) { + positionalContainer = param; + } + else if (positionalContainer == null && singleStarParam == null) { + if (param.hasDefaultValue()) { + myOptionalParams.add(param); + } + else { + myRequiredParams.add(param); + } + } + else { + if (param.hasDefaultValue()) { + myOptionalKwOnlyParams.add(param); + } + else { + myRequiredKwOnlyParams.add(param); + } } - continue; } - if (addComma) { - superCall.append(","); - } - superCall.append(param); - if (!functionParams.contains(param)) { - newFunction.append(",").append(param); - } - addComma = true; + + mySelfParam = selfParam; + myPositionalContainerParam = positionalContainer; + mySingleStarParam = singleStarParam; + myKeywordContainerParam = keywordContainer; } - for (String p : problemParams) { - newFunction.append(",").append(p); + + public boolean containsRequiredParam(@Nullable final String name) { + return ContainerUtil.exists(myRequiredParams, new Condition() { + @Override + public boolean value(PyParameter parameter) { + return name != null && name.equals(parameter.getName()); + } + }); } - if (starName != null) { - newFunction.append(",").append(starName); - if (addStar) { - if (addComma) superCall.append(","); - superCall.append(starName); - addComma = true; - } + + public boolean containsRequiredKeywordOnlyParameter(@Nullable final String name) { + return ContainerUtil.exists(myRequiredKwOnlyParams, new Condition() { + @Override + public boolean value(PyParameter parameter) { + return name != null && name.equals(parameter.getName()); + } + }); } - if (doubleStarName != null) { - newFunction.append(",").append(doubleStarName); - if (addDouble) { - if (addComma) superCall.append(","); - superCall.append(doubleStarName); - } + + @Nullable + public PyParameter getSelfParameter() { + return mySelfParam; + } + + @NotNull + public List getRequiredParameters() { + return myRequiredParams; + } + + @NotNull + public List getOptionalParameters() { + return myOptionalParams; + } + + @Nullable + public PyParameter getPositionalContainerParameter() { + return myPositionalContainerParam; + } + + @Nullable + public PyParameter getSingleStarParameter() { + return mySingleStarParam; + } + + @NotNull + public List getRequiredKeywordOnlyParameters() { + return myRequiredKwOnlyParams; + } + + @NotNull + public List getOptionalKeywordOnlyParameters() { + return myOptionalKwOnlyParams; + } + + @Nullable + public PyParameter getKeywordContainerParameter() { + return myKeywordContainerParam; } } } diff --git a/python/testData/inspections/AddCallSuperKeywordOnlyParamInInit.py b/python/testData/inspections/AddCallSuperKeywordOnlyParamInInit.py new file mode 100644 index 000000000000..e1da8962d217 --- /dev/null +++ b/python/testData/inspections/AddCallSuperKeywordOnlyParamInInit.py @@ -0,0 +1,7 @@ +class A: + def __init__(self, a): + pass + +class B(A): + def __init__(self, b, c=1, *args, kw_only): + pass \ No newline at end of file diff --git a/python/testData/inspections/AddCallSuperKeywordOnlyParamInInit_after.py b/python/testData/inspections/AddCallSuperKeywordOnlyParamInInit_after.py new file mode 100644 index 000000000000..925cac6746a3 --- /dev/null +++ b/python/testData/inspections/AddCallSuperKeywordOnlyParamInInit_after.py @@ -0,0 +1,7 @@ +class A: + def __init__(self, a): + pass + +class B(A): + def __init__(self, b, a, c=1, *args, kw_only): + super().__init__(a) \ No newline at end of file diff --git a/python/testData/inspections/AddCallSuperKeywordOnlyParamInSuperInit.py b/python/testData/inspections/AddCallSuperKeywordOnlyParamInSuperInit.py new file mode 100644 index 000000000000..5248dd92742a --- /dev/null +++ b/python/testData/inspections/AddCallSuperKeywordOnlyParamInSuperInit.py @@ -0,0 +1,8 @@ +class A: + def __init__(self, a, b=1, *args, kw_only): + pass + + +class B(A): + def __init__(self, c): + pass \ No newline at end of file diff --git a/python/testData/inspections/AddCallSuperKeywordOnlyParamInSuperInit_after.py b/python/testData/inspections/AddCallSuperKeywordOnlyParamInSuperInit_after.py new file mode 100644 index 000000000000..0ee340e47d36 --- /dev/null +++ b/python/testData/inspections/AddCallSuperKeywordOnlyParamInSuperInit_after.py @@ -0,0 +1,8 @@ +class A: + def __init__(self, a, b=1, *args, kw_only): + pass + + +class B(A): + def __init__(self, c, a, *args, kw_only): + super().__init__(a, *args, kw_only=kw_only) \ No newline at end of file diff --git a/python/testData/inspections/AddCallSuperSingleStarParamInSuperInit.py b/python/testData/inspections/AddCallSuperSingleStarParamInSuperInit.py new file mode 100644 index 000000000000..7ff489252191 --- /dev/null +++ b/python/testData/inspections/AddCallSuperSingleStarParamInSuperInit.py @@ -0,0 +1,7 @@ +class A: + def __init__(self, *, kw_only, optional_kw_only=None): + pass + +class B(A): + def __init__(self): + pass \ No newline at end of file diff --git a/python/testData/inspections/AddCallSuperSingleStarParamInSuperInitAndVarargInInit.py b/python/testData/inspections/AddCallSuperSingleStarParamInSuperInitAndVarargInInit.py new file mode 100644 index 000000000000..afca366ff1e3 --- /dev/null +++ b/python/testData/inspections/AddCallSuperSingleStarParamInSuperInitAndVarargInInit.py @@ -0,0 +1,7 @@ +class A: + def __init__(self, *, kw_only): + pass + +class B(A): + def __init__(self, *args, another_kw_only): + pass \ No newline at end of file diff --git a/python/testData/inspections/AddCallSuperSingleStarParamInSuperInitAndVarargInInit_after.py b/python/testData/inspections/AddCallSuperSingleStarParamInSuperInitAndVarargInInit_after.py new file mode 100644 index 000000000000..01cf279b21fc --- /dev/null +++ b/python/testData/inspections/AddCallSuperSingleStarParamInSuperInitAndVarargInInit_after.py @@ -0,0 +1,7 @@ +class A: + def __init__(self, *, kw_only): + pass + +class B(A): + def __init__(self, *args, another_kw_only, kw_only): + super().__init__(kw_only=kw_only) \ No newline at end of file diff --git a/python/testData/inspections/AddCallSuperSingleStarParamInSuperInit_after.py b/python/testData/inspections/AddCallSuperSingleStarParamInSuperInit_after.py new file mode 100644 index 000000000000..1eb8c4228300 --- /dev/null +++ b/python/testData/inspections/AddCallSuperSingleStarParamInSuperInit_after.py @@ -0,0 +1,7 @@ +class A: + def __init__(self, *, kw_only, optional_kw_only=None): + pass + +class B(A): + def __init__(self, *, kw_only): + super().__init__(kw_only=kw_only) \ No newline at end of file diff --git a/python/testSrc/com/jetbrains/python/PyQuickFixTest.java b/python/testSrc/com/jetbrains/python/PyQuickFixTest.java index 9a55254d8287..6ca5fe3e4211 100644 --- a/python/testSrc/com/jetbrains/python/PyQuickFixTest.java +++ b/python/testSrc/com/jetbrains/python/PyQuickFixTest.java @@ -360,6 +360,49 @@ public class PyQuickFixTest extends PyTestCase { PyBundle.message("QFIX.add.super"), true, true); } + + // PY-15867 + public void testAddCallSuperKeywordOnlyParamInSuperInit() { + runWithLanguageLevel(LanguageLevel.PYTHON30, new Runnable() { + public void run() { + doInspectionTest("AddCallSuperKeywordOnlyParamInSuperInit.py", PyMissingConstructorInspection.class, + PyBundle.message("QFIX.add.super"), true, true); + } + }); + } + + // PY-15867 + public void testAddCallSuperKeywordOnlyParamInInit() { + runWithLanguageLevel(LanguageLevel.PYTHON30, new Runnable() { + public void run() { + doInspectionTest("AddCallSuperKeywordOnlyParamInInit.py", PyMissingConstructorInspection.class, + PyBundle.message("QFIX.add.super"), true, true); + } + }); + } + + // PY-15867 + public void testAddCallSuperSingleStarParamInSuperInit() { + runWithLanguageLevel(LanguageLevel.PYTHON30, new Runnable() { + public void run() { + doInspectionTest("AddCallSuperSingleStarParamInSuperInit.py", PyMissingConstructorInspection.class, + PyBundle.message("QFIX.add.super"), true, true); + } + }); + } + + // PY-15867 + public void testAddCallSuperSingleStarParamInSuperInitAndVarargInInit() { + runWithLanguageLevel(LanguageLevel.PYTHON30, new Runnable() { + @Override + public void run() { + doInspectionTest("AddCallSuperSingleStarParamInSuperInitAndVarargInInit.py", PyMissingConstructorInspection.class, + PyBundle.message("QFIX.add.super"), true, true); + } + }); + + } + //PY-491, PY-13297 public void testAddEncoding() { doInspectionTest("AddEncoding.py", PyMandatoryEncodingInspection.class,