From 1d106de39b7f87e6f29a10f4cff73fb67f427561 Mon Sep 17 00:00:00 2001 From: Kiryl Chetyrbak Date: Wed, 19 Apr 2017 13:50:56 -0400 Subject: [PATCH] Fix 'as' statement for union types --- .../psi/impl/PyTargetExpressionImpl.java | 44 ++++++++++++------- .../com/jetbrains/python/PyTypeTest.java | 18 ++++++++ 2 files changed, 47 insertions(+), 15 deletions(-) diff --git a/python/src/com/jetbrains/python/psi/impl/PyTargetExpressionImpl.java b/python/src/com/jetbrains/python/psi/impl/PyTargetExpressionImpl.java index b682dfa45155..98b35403c3a5 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyTargetExpressionImpl.java +++ b/python/src/com/jetbrains/python/psi/impl/PyTargetExpressionImpl.java @@ -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 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; } diff --git a/python/testSrc/com/jetbrains/python/PyTypeTest.java b/python/testSrc/com/jetbrains/python/PyTypeTest.java index d7567604f38d..1c2006ee2a2d 100644 --- a/python/testSrc/com/jetbrains/python/PyTypeTest.java +++ b/python/testSrc/com/jetbrains/python/PyTypeTest.java @@ -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",