diff --git a/python/src/com/jetbrains/python/inspections/PyDataclassInspection.kt b/python/src/com/jetbrains/python/inspections/PyDataclassInspection.kt index 003617a3e718..e6ecf80c3a48 100644 --- a/python/src/com/jetbrains/python/inspections/PyDataclassInspection.kt +++ b/python/src/com/jetbrains/python/inspections/PyDataclassInspection.kt @@ -33,6 +33,10 @@ class PyDataclassInspection : PyInspection() { "attr.astuple", "attr.assoc", "attr.evolve") + + private enum class ClassOrder { + MANUALLY, DC_ORDERED, DC_UNORDERED, UNKNOWN + } } override fun buildVisitor(holder: ProblemsHolder, @@ -136,27 +140,39 @@ class PyDataclassInspection : PyInspection() { } } - override fun visitPyBinaryExpression(node: PyBinaryExpression?) { + override fun visitPyBinaryExpression(node: PyBinaryExpression) { super.visitPyBinaryExpression(node) - if (node != null && ORDER_OPERATORS.contains(node.referencedName)) { + val leftOperator = node.referencedName + if (leftOperator != null && ORDER_OPERATORS.contains(leftOperator)) { val leftClass = getInstancePyClass(node.leftExpression) ?: return val rightClass = getInstancePyClass(node.rightExpression) ?: return - val leftDataclassParameters = parseDataclassParameters(leftClass, myTypeEvalContext) + val (leftOrder, leftType) = getDataclassHierarchyOrder(leftClass, leftOperator) + if (leftOrder == ClassOrder.MANUALLY) return - if (leftClass != rightClass && - leftDataclassParameters != null && - parseDataclassParameters(rightClass, myTypeEvalContext) != null) { - registerProblem(node.psiOperator, - "'${node.referencedName}' not supported between instances of '${leftClass.name}' and '${rightClass.name}'", - ProblemHighlightType.GENERIC_ERROR) + val (rightOrder, _) = getDataclassHierarchyOrder(rightClass, PyNames.leftToRightOperatorName(leftOperator)) + + if (leftClass == rightClass) { + if (leftOrder == ClassOrder.DC_UNORDERED && rightOrder != ClassOrder.MANUALLY) { + registerProblem(node.psiOperator, + "'$leftOperator' not supported between instances of '${leftClass.name}'", + ProblemHighlightType.GENERIC_ERROR) + } } + else { + if (leftOrder == ClassOrder.DC_ORDERED || + leftOrder == ClassOrder.DC_UNORDERED || + rightOrder == ClassOrder.DC_ORDERED || + rightOrder == ClassOrder.DC_UNORDERED) { + if (leftOrder == ClassOrder.DC_ORDERED && + leftType == PyDataclassParameters.Type.ATTRS && + rightClass.isSubclass(leftClass, myTypeEvalContext)) return // attrs allows to compare ancestor and its subclass - if (leftClass == rightClass && leftDataclassParameters?.order == false && !definedReferencedOperator(leftClass, node)) { - registerProblem(node.psiOperator, - "'${node.referencedName}' not supported between instances of '${leftClass.name}'", - ProblemHighlightType.GENERIC_ERROR) + registerProblem(node.psiOperator, + "'$leftOperator' not supported between instances of '${leftClass.name}' and '${rightClass.name}'", + ProblemHighlightType.GENERIC_ERROR) + } } } } @@ -227,15 +243,31 @@ class PyDataclassInspection : PyInspection() { } } - private fun definedReferencedOperator(cls: PyClass, node: PyBinaryExpression): Boolean { - val type = cls.getType(myTypeEvalContext) ?: return false - val leftOperator = node.referencedName ?: return false + private fun getDataclassHierarchyOrder(cls: PyClass, operator: String?): Pair { + var seenUnordered: Pair? = null - val direction = AccessDirection.of(node) - if (!type.resolveMember(leftOperator, node, direction, resolveContext).isNullOrEmpty()) return true + for (current in StreamEx.of(cls).append(cls.getAncestorClasses(myTypeEvalContext))) { + val order = getDataclassOrder(current, operator) - val rightOperator = PyNames.leftToRightOperatorName(leftOperator) - return rightOperator != null && !type.resolveMember(rightOperator, node, direction, resolveContext).isNullOrEmpty() + // `order=False` just does not add comparison methods + // but it makes sense when no one in the hierarchy defines any of such methods + if (order.first == ClassOrder.DC_UNORDERED) seenUnordered = order + else if (order.first != ClassOrder.UNKNOWN) return order + } + + return if (seenUnordered != null) seenUnordered else ClassOrder.UNKNOWN to null + } + + private fun getDataclassOrder(cls: PyClass, operator: String?): Pair { + val type = cls.getType(myTypeEvalContext) + if (operator != null && + type != null && + !type.resolveMember(operator, null, AccessDirection.READ, resolveContext, false).isNullOrEmpty()) { + return ClassOrder.MANUALLY to null + } + + val parameters = parseDataclassParameters(cls, myTypeEvalContext) ?: return ClassOrder.UNKNOWN to null + return if (parameters.order) ClassOrder.DC_ORDERED to parameters.type else ClassOrder.DC_UNORDERED to parameters.type } private fun getInstancePyClass(element: PyTypedElement?): PyClass? { diff --git a/python/testData/inspections/PyDataclassInspection/comparisonForManuallyOrderedInAttrsInheritance.py b/python/testData/inspections/PyDataclassInspection/comparisonForManuallyOrderedInAttrsInheritance.py new file mode 100644 index 000000000000..9402f7f3735a --- /dev/null +++ b/python/testData/inspections/PyDataclassInspection/comparisonForManuallyOrderedInAttrsInheritance.py @@ -0,0 +1,61 @@ +from attr import dataclass + +class A7: + x: int = 1 + + def __lt__(self, other): + pass + +@dataclass(cmp=True) +class B7(A7): + y: str = "1" + +print(A7() < B7()) +print(B7() < A7()) +print(A7() < object()) +print(B7() < object()) + +class A8: + x: int = 1 + + def __lt__(self, other): + pass + +@dataclass(cmp=False) +class B8(A8): + y: str = "1" + +print(A8() < B8()) +print(B8() < A8()) +print(A8() < object()) +print(B8() < object()) + +@dataclass(cmp=True) +class A9: + x: int = 1 + +class B9(A9): + y: str = "1" + + def __lt__(self, other): + pass + +print(A9() < B9()) +print(B9() < A9()) +print(A9() < object()) +print(B9() < object()) + +@dataclass(cmp=False) +class A10: + x: int = 1 + +class B10(A10): + y: str = "1" + + def __lt__(self, other): + pass + +print(A10() < B10()) +print(B10() < A10()) +print(A10() < object()) +print(B10() < object()) \ No newline at end of file diff --git a/python/testData/inspections/PyDataclassInspection/comparisonForManuallyOrderedInStdInheritance.py b/python/testData/inspections/PyDataclassInspection/comparisonForManuallyOrderedInStdInheritance.py new file mode 100644 index 000000000000..e8ec4b7ac49a --- /dev/null +++ b/python/testData/inspections/PyDataclassInspection/comparisonForManuallyOrderedInStdInheritance.py @@ -0,0 +1,61 @@ +from dataclasses import dataclass + +class A7: + x: int = 1 + + def __lt__(self, other): + pass + +@dataclass(order=True) +class B7(A7): + y: str = "1" + +print(A7() < B7()) +print(B7() < A7()) +print(A7() < object()) +print(B7() < object()) + +class A8: + x: int = 1 + + def __lt__(self, other): + pass + +@dataclass(order=False) +class B8(A8): + y: str = "1" + +print(A8() < B8()) +print(B8() < A8()) +print(A8() < object()) +print(B8() < object()) + +@dataclass(order=True) +class A9: + x: int = 1 + +class B9(A9): + y: str = "1" + + def __lt__(self, other): + pass + +print(A9() < B9()) +print(B9() < A9()) +print(A9() < object()) +print(B9() < object()) + +@dataclass(order=False) +class A10: + x: int = 1 + +class B10(A10): + y: str = "1" + + def __lt__(self, other): + pass + +print(A10() < B10()) +print(B10() < A10()) +print(A10() < object()) +print(B10() < object()) \ No newline at end of file diff --git a/python/testData/inspections/PyDataclassInspection/comparisonInAttrsInheritance.py b/python/testData/inspections/PyDataclassInspection/comparisonInAttrsInheritance.py new file mode 100644 index 000000000000..d618080afe65 --- /dev/null +++ b/python/testData/inspections/PyDataclassInspection/comparisonInAttrsInheritance.py @@ -0,0 +1,65 @@ +from attr import dataclass + +@dataclass(cmp=True) +class A1: + x: int = 1 + +@dataclass(cmp=False) +class B1(A1): + y: str = "1" + +print(A1() < B1()) +print(B1() < B1()) + +@dataclass(cmp=False) +class A2: + x: int = 1 + +@dataclass(cmp=True) +class B2(A2): + y: str = "1" + +print(A2() < B2()) +print(B2() < B2()) + +@dataclass(cmp=False) +class A3: + x: int = 1 + +@dataclass(cmp=False) +class B3(A3): + y: str = "1" + +print(A3() < B3()) +print(B3() < B3()) + +@dataclass(cmp=True) +class A4: + x: int = 1 + +@dataclass(cmp=True) +class B4(A4): + y: str = "1" + +print(A4() < B4()) +print(B4() < B4()) + +class A5: + x: int = 1 + +@dataclass(cmp=True) +class B5(A5): + y: str = "1" + +print(A5() < B5()) +print(B5() < B5()) + +@dataclass(cmp=True) +class A6: + x: int = 1 + +class B6(A6): + y: str = "1" + +print(A6() < B6()) +print(B6() < B6()) diff --git a/python/testData/inspections/PyDataclassInspection/comparisonInStdInheritance.py b/python/testData/inspections/PyDataclassInspection/comparisonInStdInheritance.py new file mode 100644 index 000000000000..a5d36789f12f --- /dev/null +++ b/python/testData/inspections/PyDataclassInspection/comparisonInStdInheritance.py @@ -0,0 +1,65 @@ +from dataclasses import dataclass + +@dataclass(order=True) +class A1: + x: int = 1 + +@dataclass(order=False) +class B1(A1): + y: str = "1" + +print(A1() < B1()) +print(B1() < B1()) + +@dataclass(order=False) +class A2: + x: int = 1 + +@dataclass(order=True) +class B2(A2): + y: str = "1" + +print(A2() < B2()) +print(B2() < B2()) + +@dataclass(order=False) +class A3: + x: int = 1 + +@dataclass(order=False) +class B3(A3): + y: str = "1" + +print(A3() < B3()) +print(B3() < B3()) + +@dataclass(order=True) +class A4: + x: int = 1 + +@dataclass(order=True) +class B4(A4): + y: str = "1" + +print(A4() < B4()) +print(B4() < B4()) + +class A5: + x: int = 1 + +@dataclass(order=True) +class B5(A5): + y: str = "1" + +print(A5() < B5()) +print(B5() < B5()) + +@dataclass(order=True) +class A6: + x: int = 1 + +class B6(A6): + y: str = "1" + +print(A6() < B6()) +print(B6() < B6()) diff --git a/python/testSrc/com/jetbrains/python/inspections/PyDataclassInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/PyDataclassInspectionTest.java index e38f095bfdf4..61c54d6ebad7 100644 --- a/python/testSrc/com/jetbrains/python/inspections/PyDataclassInspectionTest.java +++ b/python/testSrc/com/jetbrains/python/inspections/PyDataclassInspectionTest.java @@ -117,6 +117,26 @@ public class PyDataclassInspectionTest extends PyInspectionTestCase { "print(Test > Test)"); } + // PY-28506 + public void testComparisonInStdInheritance() { + doTest(); + } + + // PY-28506 + public void testComparisonForManuallyOrderedInStdInheritance() { + doTest(); + } + + // PY-31762 + public void testComparisonInAttrsInheritance() { + doTest(); + } + + // PY-31762 + public void testComparisonForManuallyOrderedInAttrsInheritance() { + doTest(); + } + // PY-27398 public void testHelpersArgument() { doTest();