diff --git a/python/src/com/jetbrains/python/inspections/PyDataclassInspection.kt b/python/src/com/jetbrains/python/inspections/PyDataclassInspection.kt index 585a046e0a15..c04ab6876844 100644 --- a/python/src/com/jetbrains/python/inspections/PyDataclassInspection.kt +++ b/python/src/com/jetbrains/python/inspections/PyDataclassInspection.kt @@ -7,12 +7,14 @@ import com.intellij.codeInspection.LocalInspectionToolSession import com.intellij.codeInspection.ProblemHighlightType import com.intellij.codeInspection.ProblemsHolder import com.intellij.psi.PsiElementVisitor +import com.intellij.psi.PsiNameIdentifierOwner import com.intellij.util.containers.ContainerUtil import com.jetbrains.python.PyNames import com.jetbrains.python.codeInsight.stdlib.* import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider import com.jetbrains.python.psi.* import com.jetbrains.python.psi.impl.PyCallExpressionHelper +import com.jetbrains.python.psi.impl.PyEvaluator import com.jetbrains.python.psi.impl.stubs.PyDataclassFieldStubImpl import com.jetbrains.python.psi.resolve.PyResolveContext import com.jetbrains.python.psi.types.* @@ -81,6 +83,8 @@ class PyDataclassInspection : PyInspection() { } } else if (dataclassParameters.type == PyDataclassParameters.Type.ATTRS) { + processAttrsParameters(node, dataclassParameters) + node .findMethodByName(DUNDER_ATTRS_POST_INIT, false, myTypeEvalContext) ?.also { processAttrsPostInitDefinition(it, dataclassParameters) } @@ -207,21 +211,15 @@ class PyDataclassInspection : PyInspection() { var mutatingMethodsExist = false var hashMethodExists = false - cls.visitMethods( - { - when (it.name) { - "__repr__" -> reprMethodExists = true - "__eq__" -> eqMethodExists = true - in ORDER_OPERATORS -> orderMethodsExist = true - "__setattr__", "__delattr__" -> mutatingMethodsExist = true - PyNames.HASH -> hashMethodExists = true - } - - true - }, - false, - myTypeEvalContext - ) + cls.methods.forEach { + when (it.name) { + "__repr__" -> reprMethodExists = true + "__eq__" -> eqMethodExists = true + in ORDER_OPERATORS -> orderMethodsExist = true + "__setattr__", "__delattr__" -> mutatingMethodsExist = true + PyNames.HASH -> hashMethodExists = true + } + } hashMethodExists = hashMethodExists || cls.findClassAttribute(PyNames.HASH, false, myTypeEvalContext) != null @@ -256,6 +254,58 @@ class PyDataclassInspection : PyInspection() { } } + private fun processAttrsParameters(cls: PyClass, dataclassParameters: PyDataclassParameters) { + var reprMethod: PyFunction? = null + var strMethod: PyFunction? = null + val cmpMethods = mutableListOf() + val mutatingMethods = mutableListOf() + var hashMethod: PsiNameIdentifierOwner? = null + + cls.methods.forEach { + when (it.name) { + "__repr__" -> reprMethod = it + "__str__" -> strMethod = it + "__eq__", + in ORDER_OPERATORS -> cmpMethods.add(it) + "__setattr__", "__delattr__" -> mutatingMethods.add(it) + PyNames.HASH -> hashMethod = it + } + } + + hashMethod = hashMethod ?: cls.findClassAttribute(PyNames.HASH, false, myTypeEvalContext) + + // element to register problem and corresponding attr.s parameter + val problems = mutableListOf>() + + if (dataclassParameters.repr && reprMethod != null) { + problems.add(reprMethod to "repr") + } + + if (PyEvaluator.evaluateAsBoolean(PyUtil.peelArgument(dataclassParameters.others["str"]), false) && strMethod != null) { + problems.add(strMethod to "str") + } + + if (dataclassParameters.order && cmpMethods.isNotEmpty()) { + cmpMethods.forEach { problems.add(it to "cmp") } + } + + if (dataclassParameters.frozen && mutatingMethods.isNotEmpty()) { + mutatingMethods.forEach { problems.add(it to "frozen") } + } + + if (dataclassParameters.unsafeHash && hashMethod != null) { + problems.add(hashMethod to "hash") + } + + problems.forEach { + it.first?.apply { + registerProblem(nameIdentifier, + "'$name' is ignored if the class already defines '${it.second}' parameter", + ProblemHighlightType.GENERIC_ERROR_OR_WARNING) + } + } + } + private fun processDefaultFieldValue(field: PyTargetExpression) { if (field.annotationValue == null) return diff --git a/python/testData/inspections/PyDataclassInspection/attrsUselessFrozen.py b/python/testData/inspections/PyDataclassInspection/attrsUselessFrozen.py new file mode 100644 index 000000000000..d651ccea069e --- /dev/null +++ b/python/testData/inspections/PyDataclassInspection/attrsUselessFrozen.py @@ -0,0 +1,45 @@ +import attr + + +@attr.dataclass(frozen=True) +class A4: + a: int = 1 + + def __setattr__(self, key, value): + pass + +# A4(1).b = 2 + + +class Base4: + def __setattr__(self, key, value): + pass + + +@attr.dataclass(frozen=True) +class Derived4(Base4): + d: int = 1 + +# Derived4(1).b = 2 + + +@attr.dataclass(frozen=True) +class A2: + a: int = 1 + + def __delattr__(self, key): + pass + +# del A2(1).a + + +class Base2: + def __delattr__(self, key): + pass + + +@attr.dataclass(frozen=True) +class Derived2(Base2): + d: int = 1 + +# del Derived2(1).d \ No newline at end of file diff --git a/python/testData/inspections/PyDataclassInspection/attrsUselessHash.py b/python/testData/inspections/PyDataclassInspection/attrsUselessHash.py new file mode 100644 index 000000000000..1db60d8cbf39 --- /dev/null +++ b/python/testData/inspections/PyDataclassInspection/attrsUselessHash.py @@ -0,0 +1,43 @@ +import attr + + +@attr.dataclass(hash=True) +class A2: + a: int = 1 + + def __hash__(self): + pass + +print(hash(A2())) + + +class Base2: + def __hash__(self): + pass + + +@attr.dataclass(hash=True) +class Derived2(Base2): + d: int = 1 + +print(hash(Derived2())) + + +@attr.dataclass(hash=True) +class A1: + a: int = 1 + + __hash__ = None + +print(hash(A1())) + + +class Base1: + __hash__ = None + + +@attr.dataclass(hash=True) +class Derived1(Base1): + d: int = 1 + +print(hash(Derived1())) diff --git a/python/testData/inspections/PyDataclassInspection/attrsUselessOrder.py b/python/testData/inspections/PyDataclassInspection/attrsUselessOrder.py new file mode 100644 index 000000000000..b507e003ed2e --- /dev/null +++ b/python/testData/inspections/PyDataclassInspection/attrsUselessOrder.py @@ -0,0 +1,89 @@ +import attr + + +@attr.dataclass(cmp=True) +class A1: + a: int = 1 + + def __le__(self, other): + return "le1" + +print(A1(1) <= A1(2)) + + +class Base1: + def __le__(self, other): + return "le2" + + +@attr.dataclass(cmp=True) +class Derived1(Base1): + d: int = 1 + +print(Derived1(1) <= Derived1(2)) + + +@attr.dataclass(cmp=True) +class A2: + a: int = 1 + + def __lt__(self, other): + return "lt1" + +print(A2(1) < A2(2)) + + +class Base2: + def __lt__(self, other): + return "lt2" + + +@attr.dataclass(cmp=True) +class Derived2(Base2): + d: int = 1 + +print(Derived2(1) < Derived2(2)) + + +@attr.dataclass(cmp=True) +class A3: + a: int = 1 + + def __gt__(self, other): + return "gt1" + +print(A3(1) > A3(2)) + + +class Base3: + def __gt__(self, other): + return "gt2" + + +@attr.dataclass(cmp=True) +class Derived3(Base3): + d: int = 1 + +print(Derived3(1) > Derived3(2)) + + +@attr.dataclass(cmp=True) +class A4: + a: int = 1 + + def __ge__(self, other): + return "ge1" + +print(A4(1) >= A4(2)) + + +class Base4: + def __ge__(self, other): + return "ge2" + + +@attr.dataclass(cmp=True) +class Derived4(Base4): + d: int = 1 + +print(Derived4(1) >= Derived4(2)) \ No newline at end of file diff --git a/python/testData/inspections/PyDataclassInspection/attrsUselessReprStrEq.py b/python/testData/inspections/PyDataclassInspection/attrsUselessReprStrEq.py new file mode 100644 index 000000000000..42df220a875a --- /dev/null +++ b/python/testData/inspections/PyDataclassInspection/attrsUselessReprStrEq.py @@ -0,0 +1,39 @@ +import attr + + +@attr.dataclass(repr=True, cmp=True, str=True) +class A: + a: int = 1 + + def __repr__(self): + return "repr1" + + def __str__(self): + return "str1" + + def __eq__(self, other): + return "eq1" + + +class Base: + def __repr__(self): + return "repr2" + + def __str__(self): + return "str2" + + def __eq__(self, other): + return "eq2" + + +@attr.dataclass(repr=True, cmp=True, str=True) +class Derived(Base): + d: int = 1 + +print(repr(A())) +print(str(A())) +print(A() == A()) + +print(repr(Derived())) +print(str(Derived())) +print(Derived() == Derived()) diff --git a/python/testData/inspections/PyDataclassInspection/uselessFrozen.py b/python/testData/inspections/PyDataclassInspection/uselessFrozen.py index e45221788128..908bf1357b1a 100644 --- a/python/testData/inspections/PyDataclassInspection/uselessFrozen.py +++ b/python/testData/inspections/PyDataclassInspection/uselessFrozen.py @@ -14,7 +14,7 @@ class Base4: pass -@dataclasses.dataclass(order=True) +@dataclasses.dataclass(frozen=True) class Derived4(Base4): d: int = 1 @@ -32,6 +32,6 @@ class Base2: pass -@dataclasses.dataclass(order=True) +@dataclasses.dataclass(frozen=True) class Derived2(Base2): d: int = 1 \ No newline at end of file diff --git a/python/testSrc/com/jetbrains/python/inspections/PyDataclassInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/PyDataclassInspectionTest.java index e8981a19930d..9723e83cbd38 100644 --- a/python/testSrc/com/jetbrains/python/inspections/PyDataclassInspectionTest.java +++ b/python/testSrc/com/jetbrains/python/inspections/PyDataclassInspectionTest.java @@ -136,6 +136,26 @@ public class PyDataclassInspectionTest extends PyInspectionTestCase { doTest(); } + // PY-26354 + public void testAttrsUselessReprStrEq() { + doTest(); + } + + // PY-26354 + public void testAttrsUselessOrder() { + doTest(); + } + + // PY-26354 + public void testAttrsUselessFrozen() { + doTest(); + } + + // PY-26354 + public void testAttrsUselessHash() { + doTest(); + } + // PY-26354 public void testAttrsDefaultThroughKeywordAndDecorator() { doTest();