diff --git a/python/src/com/jetbrains/python/codeInsight/stdlib/PyDataclassesTypeProvider.kt b/python/src/com/jetbrains/python/codeInsight/stdlib/PyDataclassesTypeProvider.kt
index d22350b24bb2..7ab5ccfc9321 100644
--- a/python/src/com/jetbrains/python/codeInsight/stdlib/PyDataclassesTypeProvider.kt
+++ b/python/src/com/jetbrains/python/codeInsight/stdlib/PyDataclassesTypeProvider.kt
@@ -26,7 +26,7 @@ class PyDataclassesTypeProvider : PyTypeProviderBase() {
if (cls != null && name != null && parseStdDataclassParameters(cls, context)?.init == true) {
cls
.findClassAttribute(name, false, context)
- ?.let { return Ref.create(getTypeForParameter(it, context)) }
+ ?.let { return Ref.create(getTypeForParameter(it, PyDataclassParameters.Type.STD, context)) }
}
}
@@ -90,7 +90,7 @@ class PyDataclassesTypeProvider : PyTypeProviderBase() {
cls.processClassLevelDeclarations { element, _ ->
if (element is PyTargetExpression && !PyTypingTypeProvider.isClassVar(element, context)) {
- fieldToParameter(element, ellipsis, context)?.also { parameters.add(it) }
+ fieldToParameter(element, dataclassParameters.type, ellipsis, context)?.also { parameters.add(it) }
}
true
@@ -100,6 +100,7 @@ class PyDataclassesTypeProvider : PyTypeProviderBase() {
}
private fun fieldToParameter(field: PyTargetExpression,
+ dataclassType: PyDataclassParameters.Type,
ellipsis: PyNoneLiteralExpression,
context: TypeEvalContext): PyCallableParameter? {
val stub = field.stub
@@ -112,22 +113,32 @@ class PyDataclassesTypeProvider : PyTypeProviderBase() {
else -> null
}
- PyCallableParameterImpl.nonPsi(field.name, getTypeForParameter(field, context), value)
+ PyCallableParameterImpl.nonPsi(field.name, getTypeForParameter(field, dataclassType, context), value)
}
else if (!fieldStub.initValue()) {
null
}
else {
val value = if (fieldStub.hasDefault() || fieldStub.hasDefaultFactory()) ellipsis else null
- PyCallableParameterImpl.nonPsi(field.name, getTypeForParameter(field, context), value)
+ PyCallableParameterImpl.nonPsi(field.name, getTypeForParameter(field, dataclassType, context), value)
}
}
- private fun getTypeForParameter(element: PyTargetExpression, context: TypeEvalContext): PyType? {
+ private fun getTypeForParameter(element: PyTargetExpression,
+ dataclassType: PyDataclassParameters.Type,
+ context: TypeEvalContext): PyType? {
+ if (dataclassType == PyDataclassParameters.Type.ATTRS && context.maySwitchToAST(element)) {
+ (element.findAssignedValue() as? PyCallExpression)
+ ?.getKeywordArgument("type")
+ ?.let { PyTypingTypeProvider.getType(it, context) }
+ ?.apply { return get() }
+ }
+
val type = context.getType(element)
if (type is PyCollectionType && type.classQName == DATACLASSES_INITVAR_TYPE) {
return type.elementTypes.firstOrNull()
}
+
return type
}
}
\ No newline at end of file
diff --git a/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java
index 1134a469112b..1bc96d39db6a 100644
--- a/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java
+++ b/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java
@@ -297,6 +297,29 @@ public class Py3TypeCheckerInspectionTest extends PyInspectionTestCase {
runWithLanguageLevel(LanguageLevel.PYTHON37, () -> super.doMultiFileTest());
}
+ // PY-26354
+ public void testInitializingAttrs() {
+ doTestByText("import attr\n" +
+ "import typing\n" +
+ "\n" +
+ "@attr.s\n" +
+ "class Weak:\n" +
+ " x = attr.ib()\n" +
+ " y = attr.ib(default=0)\n" +
+ " z = attr.ib(default=attr.Factory(list))\n" +
+ " \n" +
+ "Weak(1, \"str\", 2)\n" +
+ "\n" +
+ "\n" +
+ "@attr.s\n" +
+ "class Strong:\n" +
+ " x = attr.ib(type=int)\n" +
+ " y = attr.ib(default=0, type=int)\n" +
+ " z = attr.ib(default=attr.Factory(list), type=typing.List[int])\n" +
+ " \n" +
+ "Strong(1, \"str\", [\"str\"])");
+ }
+
// PY-28957
public void testDataclassesReplace() {
runWithLanguageLevel(LanguageLevel.PYTHON37, () -> super.doMultiFileTest());