mirror of
https://gitflic.ru/project/openide/openide.git
synced 2026-09-27 10:03:11 +07:00
Fix 'as' statement for union types
This commit is contained in:
committed by
Andrey Vlasovskikh
parent
e25b3c2b46
commit
1d106de39b
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user