From a2e96cf59491f13daefae98e0459a4df80120fa6 Mon Sep 17 00:00:00 2001 From: Semyon Proshev Date: Tue, 24 Apr 2018 21:06:40 +0300 Subject: [PATCH] Support `type` argument of `attr.ib` (PY-26354) --- .../stdlib/PyDataclassesTypeProvider.kt | 21 +++++++++++++---- .../Py3TypeCheckerInspectionTest.java | 23 +++++++++++++++++++ 2 files changed, 39 insertions(+), 5 deletions(-) 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());