diff --git a/python/python-psi-impl/resources/messages/PyPsiBundle.properties b/python/python-psi-impl/resources/messages/PyPsiBundle.properties index 6e4447eabb2c..1a88b2624a05 100644 --- a/python/python-psi-impl/resources/messages/PyPsiBundle.properties +++ b/python/python-psi-impl/resources/messages/PyPsiBundle.properties @@ -969,6 +969,7 @@ INSP.dataclasses.method.should.be.called.on.dataclass.instances.or.types=''{0}'' INSP.dataclasses.method.should.be.called.on.dataclass.instances=''{0}'' method should be called on dataclass instances INSP.dataclasses.method.should.be.called.on.attrs.instances=''{0}'' method should be called on attrs instances INSP.dataclasses.method.should.be.called.on.attrs.types=''{0}'' method should be called on attrs types +INSP.dataclasses.expected.type.got.type.instead=Expected type ''{0}'', got ''{1}'' instead # PyDeprecationInspection INSP.NAME.deprecated.function.class.or.module=Deprecated function, class, or module diff --git a/python/python-psi-impl/src/com/jetbrains/python/inspections/PyDataclassInspection.kt b/python/python-psi-impl/src/com/jetbrains/python/inspections/PyDataclassInspection.kt index 20312f5ff0b4..349018c340b3 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/inspections/PyDataclassInspection.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/inspections/PyDataclassInspection.kt @@ -8,13 +8,14 @@ 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.intellij.util.containers.tailOrEmpty import com.jetbrains.python.PyNames import com.jetbrains.python.PyPsiBundle import com.jetbrains.python.codeInsight.* import com.jetbrains.python.codeInsight.PyDataclassNames.Attrs import com.jetbrains.python.codeInsight.PyDataclassNames.Dataclasses import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider +import com.jetbrains.python.documentation.PythonDocumentationProvider import com.jetbrains.python.psi.* import com.jetbrains.python.psi.impl.ParamHelper import com.jetbrains.python.psi.impl.PyCallExpressionHelper @@ -74,13 +75,13 @@ class PyDataclassInspection : PyInspection() { processDataclassParameters(node, dataclassParameters) val postInit = node.findMethodByName(Dataclasses.DUNDER_POST_INIT, false, myTypeEvalContext) - val localInitVars = mutableListOf() + val localInitVars = mutableListOf() node.processClassLevelDeclarations { element, _ -> if (element is PyTargetExpression) { if (!PyTypingTypeProvider.isClassVar(element, myTypeEvalContext)) { processDefaultFieldValue(element) - processAsInitVar(element, postInit)?.let { localInitVars.add(it) } + processAsInitVar(element, postInit)?.let { localInitVars.add(it.type) } } processFieldFunctionCall(node, dataclassParameters, element) @@ -518,8 +519,8 @@ class PyDataclassInspection : PyInspection() { dataclassParameters.type.asPredefinedType == PyDataclassParameters.PredefinedType.DATACLASS_TRANSFORM || PyEvaluator.evaluateAsBoolean(PyUtil.peelArgument(dataclassParameters.others["auto_attribs"]), false)) { cls.processClassLevelDeclarations { element, _ -> - if (element is PyTargetExpression - && element.annotation == null + if (element is PyTargetExpression + && element.annotation == null && resolveDataclassFieldParameters(cls, dataclassParameters, element, myTypeEvalContext) != null) { registerProblem(element, PyPsiBundle.message("INSP.dataclasses.attribute.lacks.type.annotation", element.name), ProblemHighlightType.GENERIC_ERROR) @@ -561,20 +562,23 @@ class PyDataclassInspection : PyInspection() { } } - private fun processAsInitVar(field: PyTargetExpression, postInit: PyFunction?): PyTargetExpression? { - if (isInitVar(field)) { + private fun processAsInitVar(field: PyTargetExpression, postInit: PyFunction?): InitVarField? { + val fieldType = myTypeEvalContext.getType(field) + if (isInitVar(fieldType)) { if (postInit == null) { registerProblem(field, PyPsiBundle.message("INSP.dataclasses.attribute.useless.until.post.init.declared", field.name), ProblemHighlightType.LIKE_UNUSED_SYMBOL) } - return field + return InitVarField(getInitVarType(fieldType)) } return null } + private class InitVarField(val type: PyType?) + private fun processFieldFunctionCall(dataclass: PyClass, dataclassParameters: PyDataclassParameters, field: PyTargetExpression) { val fieldStub = resolveDataclassFieldParameters(dataclass, dataclassParameters, field, myTypeEvalContext) ?: return val call = field.findAssignedValue() as? PyCallExpression ?: return @@ -595,7 +599,7 @@ class PyDataclassInspection : PyInspection() { private fun processPostInitDefinition(cls: PyClass, postInit: PyFunction, dataclassParameters: PyDataclassParameters, - localInitVars: List) { + localInitVars: List) { if (!dataclassParameters.init) { registerProblem(postInit.nameIdentifier, PyPsiBundle.message("INSP.dataclasses.post.init.would.not.be.called.until.init.parameter.set.to.true"), @@ -606,13 +610,16 @@ class PyDataclassInspection : PyInspection() { if (ParamHelper.isSelfArgsKwargsCallable(postInit, myTypeEvalContext)) return - val allInitVars = mutableListOf() + val allInitVars = mutableListOf() for (ancestor in cls.getAncestorClasses(myTypeEvalContext).asReversed()) { if (parseStdDataclassParameters(ancestor, myTypeEvalContext) == null) continue ancestor.processClassLevelDeclarations { element, _ -> - if (element is PyTargetExpression && isInitVar(element)) { - allInitVars.add(element) + if (element is PyTargetExpression) { + val fieldType = myTypeEvalContext.getType(element) + if (isInitVar(fieldType)) { + allInitVars.add(getInitVarType(fieldType)) + } } return@processClassLevelDeclarations true @@ -620,26 +627,31 @@ class PyDataclassInspection : PyInspection() { } allInitVars.addAll(localInitVars) - val implicitParameters = postInit.getParameters(myTypeEvalContext) - val parameters = if (implicitParameters.isEmpty()) emptyList() else ContainerUtil.subList(implicitParameters, 1) - - val message = if (allInitVars.size != localInitVars.size) { - PyPsiBundle.message("INSP.dataclasses.post.init.should.take.all.init.only.variables.including.inherited.in.same.order.they.defined") - } - else { - PyPsiBundle.message("INSP.dataclasses.post.init.should.take.all.init.only.variables.in.same.order.they.defined") - } - + val parameters = postInit.getParameters(myTypeEvalContext).tailOrEmpty() if (parameters.size != allInitVars.size) { + val message = if (allInitVars.size != localInitVars.size) { + PyPsiBundle.message("INSP.dataclasses.post.init.should.take.all.init.only.variables.including.inherited.in.same.order.they.defined") + } + else { + PyPsiBundle.message("INSP.dataclasses.post.init.should.take.all.init.only.variables.in.same.order.they.defined") + } registerProblem(postInit.parameterList, message, ProblemHighlightType.GENERIC_ERROR) } else { - parameters - .asSequence() - .zip(allInitVars.asSequence()) - .all { it.first.name == it.second.name } - .also { if (!it) registerProblem(postInit.parameterList, message) } + for ((index, callableParameter) in parameters.withIndex()) { + val parameter = callableParameter.parameter + if (parameter !is PyNamedParameter) continue + val annotation = PyTypingTypeProvider.getAnnotationValue(parameter, myTypeEvalContext) ?: continue + val typeFromAnnotation = PyTypingTypeProvider.getType(annotation, myTypeEvalContext) ?: continue + val initVarType = allInitVars[index] + if (!PyTypeChecker.match(typeFromAnnotation.get(), initVarType, myTypeEvalContext)) { + val initVarTypeName = PythonDocumentationProvider.getTypeName(initVarType, myTypeEvalContext) + val parameterTypeName = PythonDocumentationProvider.getVerboseTypeName(typeFromAnnotation.get(), myTypeEvalContext) + registerProblem(annotation, + PyPsiBundle.message("INSP.dataclasses.expected.type.got.type.instead", initVarTypeName, parameterTypeName)) + } + } } } @@ -695,14 +707,27 @@ class PyDataclassInspection : PyInspection() { } private fun isInitVar(field: PyTargetExpression): Boolean { - return (myTypeEvalContext.getType(field) as? PyClassType)?.classQName == Dataclasses.DATACLASSES_INITVAR + return isInitVar(myTypeEvalContext.getType(field)) } - private fun isExpectedDataclass(type: PyType?, - dataclassType: PyDataclassParameters.PredefinedType?, - allowDefinition: Boolean, - allowInstance: Boolean, - allowSubclass: Boolean): Boolean { + private fun isInitVar(fieldType: PyType?): Boolean { + return fieldType is PyCollectionType && fieldType.classQName == Dataclasses.DATACLASSES_INITVAR + } + + private fun getInitVarType(fieldType: PyType?): PyType? { + if (fieldType !is PyCollectionType || fieldType.classQName != Dataclasses.DATACLASSES_INITVAR) { + throw IllegalArgumentException() + } + return fieldType.elementTypes.singleOrNull() + } + + private fun isExpectedDataclass( + type: PyType?, + dataclassType: PyDataclassParameters.PredefinedType?, + allowDefinition: Boolean, + allowInstance: Boolean, + allowSubclass: Boolean, + ): Boolean { if (type is PyStructuralType || PyTypeChecker.isUnknown(type, myTypeEvalContext)) return true if (type is PyUnionType) return type.members.any { isExpectedDataclass(it, dataclassType, allowDefinition, allowInstance, allowSubclass) diff --git a/python/testData/inspections/PyDataclassInspection/wrongDunderPostInitSignature.py b/python/testData/inspections/PyDataclassInspection/wrongDunderPostInitSignature.py index 373fff9b2b36..e402f58b3180 100644 --- a/python/testData/inspections/PyDataclassInspection/wrongDunderPostInitSignature.py +++ b/python/testData/inspections/PyDataclassInspection/wrongDunderPostInitSignature.py @@ -15,5 +15,5 @@ class B: b: dataclasses.InitVar[str] c: dataclasses.InitVar[bytes] - def __post_init__(self, c, b): + def __post_init__(self, c: bytes, b: bytes): pass \ No newline at end of file