diff --git a/python/src/com/jetbrains/python/inspections/PyDataclassInspection.kt b/python/src/com/jetbrains/python/inspections/PyDataclassInspection.kt index 38b432b5c95b..ca15ba6bc1e3 100644 --- a/python/src/com/jetbrains/python/inspections/PyDataclassInspection.kt +++ b/python/src/com/jetbrains/python/inspections/PyDataclassInspection.kt @@ -22,6 +22,8 @@ class PyDataclassInspection : PyInspection() { companion object { private val ORDER_OPERATORS = setOf("__lt__", "__le__", "__gt__", "__ge__") private val DATACLASSES_HELPERS = setOf("dataclasses.fields", "dataclasses.asdict", "dataclasses.astuple", "dataclasses.replace") + private val ATTRS_HELPERS = + setOf("attr.__init__.fields", "attr.__init__.asdict", "attr.__init__.astuple", "attr.__init__.assoc", "attr.__init__.evolve") } override fun buildVisitor(holder: ProblemsHolder, @@ -131,13 +133,24 @@ class PyDataclassInspection : PyInspection() { val callee = markedCallee?.element val calleeQName = callee?.qualifiedName - if (markedCallee != null && callee != null && DATACLASSES_HELPERS.contains(calleeQName)) { + if (markedCallee != null && callee != null) { + val dataclassType = when { + DATACLASSES_HELPERS.contains(calleeQName) -> PyDataclassParameters.Type.STD + ATTRS_HELPERS.contains(calleeQName) -> PyDataclassParameters.Type.ATTRS + else -> return + } + val mapping = PyCallExpressionHelper.mapArguments(node, markedCallee, myTypeEvalContext) val dataclassParameter = callee.getParameters(myTypeEvalContext).firstOrNull() val dataclassArgument = mapping.mappedParameters.entries.firstOrNull { it.value == dataclassParameter }?.key - processHelperDataclassArgument(dataclassArgument, calleeQName!!) + if (dataclassType == PyDataclassParameters.Type.STD) { + processHelperDataclassArgument(dataclassArgument, calleeQName!!) + } + else if (dataclassType == PyDataclassParameters.Type.ATTRS) { + processHelperAttrsArgument(dataclassArgument, calleeQName!!) + } } } } @@ -306,24 +319,41 @@ class PyDataclassInspection : PyInspection() { val allowDefinition = calleeQName == "dataclasses.fields" - if (isNotDataclass(myTypeEvalContext.getType(argument), allowDefinition)) { + if (isNotExpectedDataclass(myTypeEvalContext.getType(argument), PyDataclassParameters.Type.STD, allowDefinition, true)) { val message = "'$calleeQName' method should be called on dataclass instances" + if (allowDefinition) " or types" else "" registerProblem(argument, message) } } + private fun processHelperAttrsArgument(argument: PyExpression?, calleeQName: String) { + if (argument == null) return + + val instance = calleeQName != "attr.__init__.fields" + + if (isNotExpectedDataclass(myTypeEvalContext.getType(argument), PyDataclassParameters.Type.ATTRS, !instance, instance)) { + val presentableCalleeQName = calleeQName.replaceFirst(".__init__.", ".") + val message = "'$presentableCalleeQName' method should be called on attrs " + if (instance) "instances" else "types" + + registerProblem(argument, message) + } + } + private fun isInitVar(field: PyTargetExpression): Boolean { return (myTypeEvalContext.getType(field) as? PyClassType)?.classQName == DATACLASSES_INITVAR_TYPE } - private fun isNotDataclass(type: PyType?, allowDefinition: Boolean): Boolean { + private fun isNotExpectedDataclass(type: PyType?, + dataclassType: PyDataclassParameters.Type, + allowDefinition: Boolean, + allowInstance: Boolean): Boolean { if (type is PyStructuralType || PyTypeChecker.isUnknown(type, myTypeEvalContext)) return false - if (type is PyUnionType) return type.members.all { isNotDataclass(it, allowDefinition) } + if (type is PyUnionType) return type.members.all { isNotExpectedDataclass(it, dataclassType, allowDefinition, allowInstance) } return type !is PyClassType || !allowDefinition && type.isDefinition || - parseStdDataclassParameters(type.pyClass, myTypeEvalContext) == null + !allowInstance && !type.isDefinition || + parseDataclassParameters(type.pyClass, myTypeEvalContext)?.type != dataclassType } } } \ No newline at end of file diff --git a/python/testData/inspections/PyDataclassInspection/attrsHelpersArgument.py b/python/testData/inspections/PyDataclassInspection/attrsHelpersArgument.py new file mode 100644 index 000000000000..d425eb835f2d --- /dev/null +++ b/python/testData/inspections/PyDataclassInspection/attrsHelpersArgument.py @@ -0,0 +1,69 @@ +import attr +from typing import Type, Union + + +class A: + pass + + +attr.fields(A) +attr.fields(A()) + +attr.asdict(A()) +attr.astuple(A()) +attr.assoc(A()) +attr.evolve(A()) + + +@attr.s +class B: + pass + + +attr.fields(B) +attr.fields(B()) + +attr.asdict(B()) +attr.astuple(B()) +attr.assoc(B()) +attr.evolve(B()) + +attr.asdict(B) +attr.astuple(B) +attr.assoc(B) +attr.evolve(B) + + +def unknown(p): + attr.fields(p) + + attr.asdict(p) + attr.astuple(p) + + +def structural(p): + print(len(p)) + attr.fields(p) + + attr.asdict(p) + attr.astuple(p) + attr.assoc(p) + attr.evolve(p) + + +def union1(p: Union[A, B]): + attr.fields(p) + + attr.asdict(p) + attr.astuple(p) + attr.assoc(p) + attr.evolve(p) + + +def union2(p: Union[Type[A], Type[B]]): + attr.fields(p) + + attr.asdict(p) + attr.astuple(p) + attr.assoc(p) + attr.evolve(p) \ 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 cb227f6a7b9f..6eb91e563bad 100644 --- a/python/testSrc/com/jetbrains/python/inspections/PyDataclassInspectionTest.java +++ b/python/testSrc/com/jetbrains/python/inspections/PyDataclassInspectionTest.java @@ -76,6 +76,11 @@ public class PyDataclassInspectionTest extends PyInspectionTestCase { doTest(); } + // PY-26354 + public void testAttrsHelpersArgument() { + doTest(); + } + // PY-27398 public void testAccessToInitVar() { doTest();