Support type argument of attr.ib (PY-26354)

This commit is contained in:
Semyon Proshev
2018-05-29 16:49:52 +03:00
parent 7261e9121a
commit a2e96cf594
2 changed files with 39 additions and 5 deletions
@@ -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
}
}
@@ -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, <warning descr=\"Expected type 'int', got 'str' instead\">\"str\"</warning>, <warning descr=\"Expected type 'List[int]', got 'List[str]' instead\">[\"str\"]</warning>)");
}
// PY-28957
public void testDataclassesReplace() {
runWithLanguageLevel(LanguageLevel.PYTHON37, () -> super.doMultiFileTest());