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();