Infer 'x: Optional[T]' from 'x: T = None' according to PEP 484 (PY-15206)

This commit is contained in:
Andrey Vlasovskikh
2015-04-01 16:06:32 +03:00
parent ba7d99f478
commit c832e56c07
2 changed files with 25 additions and 3 deletions
@@ -51,6 +51,7 @@ public class PyTypingTypeProvider extends PyTypeProviderBase {
.add("typing.Protocol")
.build();
@Nullable
public Ref<PyType> getParameterType(@NotNull PyNamedParameter param, @NotNull PyFunction func, @NotNull TypeEvalContext context) {
final PyAnnotation annotation = param.getAnnotation();
if (annotation != null) {
@@ -59,7 +60,8 @@ public class PyTypingTypeProvider extends PyTypeProviderBase {
if (value != null) {
final PyType type = getTypingType(value, context);
if (type != null) {
return Ref.create(type);
final PyType optionalType = getOptionalTypeFromDefaultNone(param, type, context);
return Ref.create(optionalType != null ? optionalType : type);
}
}
}
@@ -92,6 +94,20 @@ public class PyTypingTypeProvider extends PyTypeProviderBase {
return type instanceof PyClassType && "typing.Any".equals(((PyClassType)type).getPyClass().getQualifiedName());
}
@Nullable
private static PyType getOptionalTypeFromDefaultNone(@NotNull PyNamedParameter param,
@NotNull PyType type,
@NotNull TypeEvalContext context) {
final PyExpression defaultValue = param.getDefaultValue();
if (defaultValue != null) {
final PyType defaultType = context.getType(defaultValue);
if (defaultType instanceof PyNoneType) {
return PyUnionType.union(type, defaultType);
}
}
return null;
}
@Nullable
private static PyType getGenericConstructorType(@NotNull PyFunction function, @NotNull TypeEvalContext context) {
if (PyUtil.isInit(function)) {
@@ -152,7 +168,7 @@ public class PyTypingTypeProvider extends PyTypeProviderBase {
if (unionType != null) {
return unionType;
}
final Ref<PyType> optionalType = getOptionalType(expression, context);
final Ref<PyType> optionalType = getOptionalTypeFromDefaultNone(expression, context);
if (optionalType != null) {
return optionalType.get();
}
@@ -203,7 +219,7 @@ public class PyTypingTypeProvider extends PyTypeProviderBase {
}
@Nullable
private static Ref<PyType> getOptionalType(@NotNull PyExpression expression, @NotNull TypeEvalContext context) {
private static Ref<PyType> getOptionalTypeFromDefaultNone(@NotNull PyExpression expression, @NotNull TypeEvalContext context) {
if (expression instanceof PySubscriptionExpression) {
final PySubscriptionExpression subscriptionExpr = (PySubscriptionExpression)expression;
final PyExpression operand = subscriptionExpr.getOperand();
@@ -257,6 +257,12 @@ public class PyTypingTest extends PyTestCase {
" pass\n");
}
public void testOptionalFromDefaultNone() {
doTest("Optional[int]",
"def foo(expr: int = None):\n" +
" pass\n");
}
private void doTest(@NotNull String expectedType, @NotNull String text) {
myFixture.copyDirectoryToProject("typing", "");
myFixture.configureByText(PythonFileType.INSTANCE, text);