From a6c3621086fc7f13c7c1dc486f80a63440b5dd8c Mon Sep 17 00:00:00 2001 From: Semyon Proshev Date: Mon, 2 Apr 2018 18:21:46 +0300 Subject: [PATCH] Reduce code duplication in determining if default argument value is mutable (PY-27517) --- .../inspections/PyDataclassInspection.kt | 16 +++--------- .../PyDefaultArgumentInspection.java | 17 +++++-------- .../src/com/jetbrains/python/psi/PyUtil.java | 19 ++++++++++++++ .../defaultFieldValue.py | 14 +++++------ .../PyDefaultArgumentInspection/expected.xml | 8 ------ .../PyDefaultArgumentInspection/src/test.py | 3 --- .../PyDefaultArgumentInspection/test.py | 25 +++++++++++++++++++ .../python/PythonInspectionsTest.java | 3 +-- 8 files changed, 62 insertions(+), 43 deletions(-) delete mode 100644 python/testData/inspections/PyDefaultArgumentInspection/expected.xml delete mode 100644 python/testData/inspections/PyDefaultArgumentInspection/src/test.py create mode 100644 python/testData/inspections/PyDefaultArgumentInspection/test.py diff --git a/python/src/com/jetbrains/python/inspections/PyDataclassInspection.kt b/python/src/com/jetbrains/python/inspections/PyDataclassInspection.kt index 27472a3ea8e1..6297e6e2ecc8 100644 --- a/python/src/com/jetbrains/python/inspections/PyDataclassInspection.kt +++ b/python/src/com/jetbrains/python/inspections/PyDataclassInspection.kt @@ -15,7 +15,6 @@ import com.jetbrains.python.codeInsight.stdlib.DataclassParameters import com.jetbrains.python.codeInsight.stdlib.parseDataclassParameters import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider import com.jetbrains.python.psi.* -import com.jetbrains.python.psi.impl.PyBuiltinCache import com.jetbrains.python.psi.impl.PyCallExpressionHelper import com.jetbrains.python.psi.impl.stubs.PyDataclassFieldStubImpl import com.jetbrains.python.psi.resolve.PyResolveContext @@ -242,17 +241,10 @@ class PyDataclassInspection : PyInspection() { if (field.annotationValue == null) return val value = field.findAssignedValue() - val valueClass = getInstancePyClass(value) - - if (valueClass != null) { - val builtinCache = PyBuiltinCache.getInstance(field) - val disallowed = setOf(builtinCache.listType?.pyClass, builtinCache.setType?.pyClass, builtinCache.dictType?.pyClass) - - if (valueClass in disallowed || valueClass.getAncestorClasses(myTypeEvalContext).find(disallowed::contains) != null) { - registerProblem(value, - "Mutable default '${valueClass.name}' is not allowed. Use 'default_factory'", - ProblemHighlightType.GENERIC_ERROR) - } + if (PyUtil.isForbiddenMutableDefault(value, myTypeEvalContext)) { + registerProblem(value, + "Mutable default '${value?.text}' is not allowed. Use 'default_factory'", + ProblemHighlightType.GENERIC_ERROR) } } diff --git a/python/src/com/jetbrains/python/inspections/PyDefaultArgumentInspection.java b/python/src/com/jetbrains/python/inspections/PyDefaultArgumentInspection.java index 4feb3e272539..984bd1f2cd2d 100644 --- a/python/src/com/jetbrains/python/inspections/PyDefaultArgumentInspection.java +++ b/python/src/com/jetbrains/python/inspections/PyDefaultArgumentInspection.java @@ -20,7 +20,9 @@ import com.intellij.codeInspection.ProblemsHolder; import com.intellij.psi.PsiElementVisitor; import com.jetbrains.python.PyBundle; import com.jetbrains.python.inspections.quickfix.PyDefaultArgumentQuickFix; -import com.jetbrains.python.psi.*; +import com.jetbrains.python.psi.PyExpression; +import com.jetbrains.python.psi.PyNamedParameter; +import com.jetbrains.python.psi.PyUtil; import org.jetbrains.annotations.Nls; import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; @@ -51,16 +53,9 @@ public class PyDefaultArgumentInspection extends PyInspection { @Override public void visitPyNamedParameter(PyNamedParameter node) { - PyExpression defaultValue = node.getDefaultValue(); - if (defaultValue != null) { - if (defaultValue instanceof PyListLiteralExpression || defaultValue instanceof PyDictLiteralExpression) { - registerProblem(defaultValue, "Default argument value is mutable", new PyDefaultArgumentQuickFix()); - } - if (defaultValue instanceof PyCallExpression) { - PyExpression callee = ((PyCallExpression)defaultValue).getCallee(); - if (callee != null && "dict".equals(callee.getText())) - registerProblem(defaultValue, "Default argument value is mutable", new PyDefaultArgumentQuickFix()); - } + final PyExpression defaultValue = node.getDefaultValue(); + if (PyUtil.isForbiddenMutableDefault(defaultValue, myTypeEvalContext)) { + registerProblem(defaultValue, "Default argument value is mutable", new PyDefaultArgumentQuickFix()); } } } diff --git a/python/src/com/jetbrains/python/psi/PyUtil.java b/python/src/com/jetbrains/python/psi/PyUtil.java index b7e6f18167c2..1ffaadc9025b 100644 --- a/python/src/com/jetbrains/python/psi/PyUtil.java +++ b/python/src/com/jetbrains/python/psi/PyUtil.java @@ -1872,6 +1872,25 @@ public class PyUtil { return loop; } + public static boolean isForbiddenMutableDefault(@Nullable PyTypedElement value, @NotNull TypeEvalContext context) { + if (value == null) return false; + + final PyClassType type = as(context.getType(value), PyClassType.class); + if (type != null && !type.isDefinition()) { + final PyBuiltinCache builtinCache = PyBuiltinCache.getInstance(value); + final Set forbiddenClasses = StreamEx + .of(builtinCache.getListType(), builtinCache.getSetType(), builtinCache.getDictType()) + .nonNull() + .map(PyClassType::getPyClass) + .toSet(); + + final PyClass cls = type.getPyClass(); + return forbiddenClasses.contains(cls) || ContainerUtil.exists(cls.getAncestorClasses(context), forbiddenClasses::contains); + } + + return false; + } + /** * This helper class allows to collect various information about AST nodes composing {@link PyStringLiteralExpression}. */ diff --git a/python/testData/inspections/PyDataclassInspection/defaultFieldValue.py b/python/testData/inspections/PyDataclassInspection/defaultFieldValue.py index e4846944169b..57c633d7b925 100644 --- a/python/testData/inspections/PyDataclassInspection/defaultFieldValue.py +++ b/python/testData/inspections/PyDataclassInspection/defaultFieldValue.py @@ -5,19 +5,19 @@ from collections import OrderedDict @dataclasses.dataclass class A: - a: List[int] = [] - b: List[int] = list() - c: Set[int] = {1} - d: Set[int] = set() + a: List[int] = [] + b: List[int] = list() + c: Set[int] = {1} + d: Set[int] = set() e: Tuple[int, ...] = () f: Tuple[int, ...] = tuple() g: ClassVar[List[int]] = [] h: ClassVar = [] - i: Dict[int, int] = {1: 2} - j: Dict[int, int] = dict() + i: Dict[int, int] = {1: 2} + j: Dict[int, int] = dict() k = [] l = list() - m: Dict[int, int] = OrderedDict() + m: Dict[int, int] = OrderedDict() n: FrozenSet[int] = frozenset() a2: Type[List[int]] = list b2: Type[Set[int]] = set diff --git a/python/testData/inspections/PyDefaultArgumentInspection/expected.xml b/python/testData/inspections/PyDefaultArgumentInspection/expected.xml deleted file mode 100644 index 17dd5717ec83..000000000000 --- a/python/testData/inspections/PyDefaultArgumentInspection/expected.xml +++ /dev/null @@ -1,8 +0,0 @@ - - - - test.py - 1 - Default argument value is mutable - - \ No newline at end of file diff --git a/python/testData/inspections/PyDefaultArgumentInspection/src/test.py b/python/testData/inspections/PyDefaultArgumentInspection/src/test.py deleted file mode 100644 index 0c8594040ec9..000000000000 --- a/python/testData/inspections/PyDefaultArgumentInspection/src/test.py +++ /dev/null @@ -1,3 +0,0 @@ -def f(a, L=[]): - L.append(a) - return L diff --git a/python/testData/inspections/PyDefaultArgumentInspection/test.py b/python/testData/inspections/PyDefaultArgumentInspection/test.py new file mode 100644 index 000000000000..beadaa0c5c71 --- /dev/null +++ b/python/testData/inspections/PyDefaultArgumentInspection/test.py @@ -0,0 +1,25 @@ +def f(a, L=[]): + L.append(a) + return L + +def f(a, L=list()): + L.append(a) + return L + + +def f(a, L=set()): + L.append(a) + return L + +def f(a, L={}): + L.append(a) + return L + + +def f(a, L=dict()): + L.append(a) + return L + +def f(a, L={1: 2}): + L.append(a) + return L diff --git a/python/testSrc/com/jetbrains/python/PythonInspectionsTest.java b/python/testSrc/com/jetbrains/python/PythonInspectionsTest.java index 6065ce4f9e78..c6f7f1afe0f3 100644 --- a/python/testSrc/com/jetbrains/python/PythonInspectionsTest.java +++ b/python/testSrc/com/jetbrains/python/PythonInspectionsTest.java @@ -90,8 +90,7 @@ public class PythonInspectionsTest extends PyTestCase { } public void testPyDefaultArgumentInspection() { - LocalInspectionTool inspection = new PyDefaultArgumentInspection(); - doTest(getTestName(false), inspection); + doHighlightingTest(PyDefaultArgumentInspection.class); } public void testPyDocstringInspection() {