Don't wrap types returned by type providers in Coroutine and AsyncGenerator

It's their responsibility to provide the final return type for a function.
Only inferred (i.e. when there's no additional type hints) return value
types should be postprocessed like that in order to get not the
immediate type of values in return statements but the properly
parametrized special type from typing for genetators and coroutines.

Otherwise, we might end up returning something like
Coroutine[Coroutine[int]] for coroutines annotated using .pyi stubs,
because we don't take into account that PyiTypeProvider already
wraps the return type from the stub's annotation in typing.Coroutine.
This commit is contained in:
Mikhail Golubev
2017-02-03 17:06:23 +03:00
parent 10ac9a4b38
commit aa41ca3616
5 changed files with 47 additions and 27 deletions
@@ -26,6 +26,7 @@ import com.intellij.psi.PsiPolyVariantReference;
import com.intellij.psi.util.CachedValueProvider;
import com.intellij.psi.util.CachedValuesManager;
import com.intellij.psi.util.PsiTreeUtil;
import com.intellij.psi.util.QualifiedName;
import com.intellij.util.containers.ContainerUtil;
import com.intellij.util.containers.HashSet;
import com.jetbrains.python.PyCustomType;
@@ -36,6 +37,7 @@ import com.jetbrains.python.codeInsight.functionTypeComments.psi.PyParameterType
import com.jetbrains.python.psi.*;
import com.jetbrains.python.psi.impl.PyExpressionCodeFragmentImpl;
import com.jetbrains.python.psi.resolve.PyResolveContext;
import com.jetbrains.python.psi.resolve.PyResolveImportUtil;
import com.jetbrains.python.psi.types.*;
import one.util.streamex.StreamEx;
import org.jetbrains.annotations.NotNull;
@@ -239,11 +241,15 @@ public class PyTypingTypeProvider extends PyTypeProviderBase {
final PyExpression value = getReturnTypeAnnotation(function);
if (value != null) {
final Ref<PyType> typeRef = getType(value, new Context(context));
final PyType type = Ref.deref(typeRef);
if (isInit && type instanceof PyNoneType) {
return null;
if (typeRef != null) {
if (isInit && typeRef.get() instanceof PyNoneType) {
return null;
}
if (function.isAsync() && function.isAsyncAllowed() && !function.isGenerator()) {
return Ref.create(wrapInCoroutineType(typeRef.get(), callable));
}
return typeRef;
}
return typeRef;
}
}
return null;
@@ -716,6 +722,13 @@ public class PyTypingTypeProvider extends PyTypeProviderBase {
return null;
}
@Nullable
public static PyType wrapInCoroutineType(@Nullable PyType returnType, @NotNull PsiElement resolveAnchor) {
final PyClass coroutine = as(PyResolveImportUtil.resolveTopLevelMember(QualifiedName.fromDottedString(COROUTINE),
PyResolveImportUtil.fromFoothold(resolveAnchor)), PyClass.class);
return coroutine != null ? new PyCollectionTypeImpl(coroutine, false, Arrays.asList(null, null, returnType)) : null;
}
private static class Context {
@NotNull private final TypeEvalContext myContext;
@NotNull private final Set<PsiElement> myCache = new HashSet<>();
@@ -38,10 +38,10 @@ import com.intellij.util.containers.JBIterable;
import com.jetbrains.python.PyElementTypes;
import com.jetbrains.python.PyNames;
import com.jetbrains.python.PyTokenTypes;
import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider;
import com.jetbrains.python.codeInsight.controlflow.ControlFlowCache;
import com.jetbrains.python.codeInsight.controlflow.ScopeOwner;
import com.jetbrains.python.codeInsight.dataflow.scope.ScopeUtil;
import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider;
import com.jetbrains.python.documentation.docstrings.DocStringUtil;
import com.jetbrains.python.psi.*;
import com.jetbrains.python.psi.resolve.PyResolveImportUtil;
@@ -187,12 +187,6 @@ public class PyFunctionImpl extends PyBaseElementImpl<PyFunctionStub> implements
@Nullable
@Override
public PyType getReturnType(@NotNull TypeEvalContext context, @NotNull TypeEvalContext.Key key) {
final PyType type = getReturnType(context);
return isAsync() && isAsyncAllowed() ? createCoroutineType(type) : type;
}
@Nullable
private PyType getReturnType(@NotNull TypeEvalContext context) {
for (PyTypeProvider typeProvider : Extensions.getExtensions(PyTypeProvider.EP_NAME)) {
final Ref<PyType> returnTypeRef = typeProvider.getReturnType(this, context);
if (returnTypeRef != null) {
@@ -200,15 +194,21 @@ public class PyFunctionImpl extends PyBaseElementImpl<PyFunctionStub> implements
}
}
PyType inferredType = null;
if (context.allowReturnTypes(this)) {
final Ref<? extends PyType> yieldTypeRef = getYieldStatementType(context);
if (yieldTypeRef != null) {
return yieldTypeRef.get();
inferredType = yieldTypeRef.get();
}
return getReturnStatementType(context);
else {
inferredType = getReturnStatementType(context);
}
}
return null;
if (isAsync() && isAsyncAllowed()) {
return createAsyncType(inferredType);
}
return inferredType;
}
@Nullable
@@ -383,12 +383,11 @@ public class PyFunctionImpl extends PyBaseElementImpl<PyFunctionStub> implements
}
@Nullable
private PyType createCoroutineType(@Nullable PyType returnType) {
if (isGenerator()) {
// Re-wrap typing.Generator[Y, S, R] into typing.AsyncGenerator[Y, Any]
if (returnType instanceof PyCollectionType) {
final PyClassLikeType classType = as(returnType, PyClassLikeType.class);
if (classType != null && PyTypingTypeProvider.GENERATOR.equals(classType.getClassQName())) {
private PyType createAsyncType(@Nullable PyType returnType) {
if (returnType instanceof PyCollectionType) {
final PyClassLikeType classType = as(returnType, PyClassLikeType.class);
if (classType != null) {
if (PyTypingTypeProvider.GENERATOR.equals(classType.getClassQName())) {
final QualifiedName asyncGeneratorName = QualifiedName.fromDottedString(PyTypingTypeProvider.ASYNC_GENERATOR);
final PsiElement resolvedGenerator = PyResolveImportUtil.resolveTopLevelMember(asyncGeneratorName,
PyResolveImportUtil.fromFoothold(this));
@@ -397,16 +396,14 @@ public class PyFunctionImpl extends PyBaseElementImpl<PyFunctionStub> implements
return new PyCollectionTypeImpl(asyncGenerator, false,
Arrays.asList(((PyCollectionType)returnType).getIteratedItemType(), null));
}
else {
return null;
}
}
}
// Leave the type as is
return returnType;
}
else {
final PyClass coroutine = as(PyResolveImportUtil.resolveTopLevelMember(QualifiedName.fromDottedString(PyTypingTypeProvider.COROUTINE),
PyResolveImportUtil.fromFoothold(this)), PyClass.class);
return coroutine != null ? new PyCollectionTypeImpl(coroutine, false, Arrays.asList(null, null, returnType)) : null;
}
return PyTypingTypeProvider.wrapInCoroutineType(returnType, this);
}
public PyFunction asMethod() {
@@ -0,0 +1,4 @@
async def f():
return 42
co<caret>routine = f()
@@ -0,0 +1,2 @@
async def f() -> int:
...
@@ -101,6 +101,10 @@ public class PyiTypeTest extends PyTestCase {
doTest("int");
}
public void testCoroutineType() {
doTest("Coroutine[Any, Any, int]");
}
public void testPyiOnPythonPath() {
addPyiStubsToContentRoot(myFixture);
doTest("int");