Fix 'as' statement for union types

This commit is contained in:
Kiryl Chetyrbak
2017-07-03 22:06:38 +03:00
committed by Andrey Vlasovskikh
parent e25b3c2b46
commit 1d106de39b
2 changed files with 47 additions and 15 deletions
@@ -57,11 +57,13 @@ import com.jetbrains.python.psi.stubs.PyClassStub;
import com.jetbrains.python.psi.stubs.PyFunctionStub;
import com.jetbrains.python.psi.stubs.PyTargetExpressionStub;
import com.jetbrains.python.psi.types.*;
import one.util.streamex.StreamEx;
import org.jetbrains.annotations.NotNull;
import org.jetbrains.annotations.Nullable;
import javax.swing.*;
import java.util.*;
import java.util.stream.Collectors;
import static com.jetbrains.python.psi.PyUtil.as;
@@ -244,23 +246,35 @@ public class PyTargetExpressionImpl extends PyBaseElementImpl<PyTargetExpression
if (expression != null) {
final PyType exprType = context.getType(expression);
if (exprType instanceof PyClassType) {
final PyClass cls = ((PyClassType)exprType).getPyClass();
final PyFunction enter = cls.findMethodByName(PyNames.ENTER, true, null);
if (enter != null) {
final PyType enterType = enter.getCallType(expression, Collections.emptyMap(), context);
if (enterType != null) {
return enterType;
}
for (PyTypeProvider provider : Extensions.getExtensions(PyTypeProvider.EP_NAME)) {
PyType typeFromProvider = provider.getContextManagerVariableType(cls, expression, context);
if (typeFromProvider != null) {
return typeFromProvider;
}
}
// Guess the return type of __enter__
return PyUnionType.createWeakType(exprType);
return getEnterTypeFromPyClass(context, expression, (PyClassType)exprType);
}
else if (exprType instanceof PyUnionType) {
List<PyType> collect = StreamEx.of(((PyUnionType)exprType).getMembers())
.select(PyClassType.class)
.map(t -> getEnterTypeFromPyClass(context, expression, t))
.toList();
return PyUnionType.union(collect);
}
}
return null;
}
private static PyType getEnterTypeFromPyClass(TypeEvalContext context, PyExpression expression, @NotNull PyClassType exprType) {
final PyClass cls = exprType.getPyClass();
final PyFunction enter = cls.findMethodByName(PyNames.ENTER, true, null);
if (enter != null) {
final PyType enterType = enter.getCallType(expression, Collections.emptyMap(), context);
if (enterType != null) {
return enterType;
}
for (PyTypeProvider provider : Extensions.getExtensions(PyTypeProvider.EP_NAME)) {
PyType typeFromProvider = provider.getContextManagerVariableType(cls, expression, context);
if (typeFromProvider != null) {
return typeFromProvider;
}
}
// Guess the return type of __enter__
return PyUnionType.createWeakType(exprType);
}
return null;
}
@@ -1702,6 +1702,24 @@ public class PyTypeTest extends PyTestCase {
"expr = max(l)");
}
public void testWithAsType() {
doTest("Union[A, B]",
"from typing import Union\n" +
"\n" +
"class A(object):\n" +
" def __enter__(self):\n" +
" return self\n" +
"\n" +
"class B(object):\n" +
" def __enter__(self):\n" +
" return self\n" +
"\n" +
"def f(x):\n" +
" # type: (Union[A, B]) -> None\n" +
" with x as expr:\n" +
" pass");
}
// PY-23634
public void testMinListKnownElements() {
doTest("int",