Fix inferring type for await on imported coroutines (PY-26847)

This commit is contained in:
Semyon Proshev
2017-11-11 00:00:32 +03:00
parent dc41db37d9
commit 313daaaba9
3 changed files with 34 additions and 39 deletions
@@ -1,18 +1,4 @@
/*
* Copyright 2000-2017 JetBrains s.r.o.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
// Copyright 2000-2017 JetBrains s.r.o. Use of this source code is governed by the Apache 2.0 license that can be found in the LICENSE file.
package com.jetbrains.python.psi.impl;
import com.intellij.lang.ASTNode;
@@ -90,9 +76,9 @@ public class PyPrefixExpressionImpl extends PyElementImpl implements PyPrefixExp
final PyExpression operand = getOperand();
if (operand != null) {
final PyType operandType = context.getType(operand);
final PyType type = getGeneratorReturnType(operandType, context);
final Ref<PyType> type = getGeneratorReturnType(operandType, context);
if (type != null) {
return type;
return type.get();
}
}
}
@@ -101,7 +87,7 @@ public class PyPrefixExpressionImpl extends PyElementImpl implements PyPrefixExp
if (resolved instanceof PyCallable) {
// TODO: Make PyPrefixExpression a PyCallSiteExpression, use getCallType() here and analyze it in PyTypeChecker.analyzeCallSite()
final PyType returnType = ((PyCallable)resolved).getReturnType(context, key);
return isAwait ? getGeneratorReturnType(returnType, context) : returnType;
return isAwait ? Ref.deref(getGeneratorReturnType(returnType, context)) : returnType;
}
return null;
}
@@ -124,7 +110,7 @@ public class PyPrefixExpressionImpl extends PyElementImpl implements PyPrefixExp
@Override
public String getReferencedName() {
PyElementType t = getOperator();
final PyElementType t = getOperator();
if (t == PyTokenTypes.PLUS) {
return PyNames.POS;
}
@@ -141,22 +127,22 @@ public class PyPrefixExpressionImpl extends PyElementImpl implements PyPrefixExp
}
@Nullable
private static PyType getGeneratorReturnType(@Nullable PyType type, @NotNull TypeEvalContext context) {
private static Ref<PyType> getGeneratorReturnType(@Nullable PyType type, @NotNull TypeEvalContext context) {
if (type instanceof PyClassLikeType && type instanceof PyCollectionType) {
if (type instanceof PyClassType && PyNames.AWAITABLE.equals(((PyClassType)type).getPyClass().getName())) {
return ((PyCollectionType)type).getIteratedItemType();
return Ref.create(((PyCollectionType)type).getIteratedItemType());
}
else {
return Ref.deref(PyTypingTypeProvider.coroutineOrGeneratorElementType(type, context));
return PyTypingTypeProvider.coroutineOrGeneratorElementType(type, context);
}
}
else if (type instanceof PyUnionType) {
final List<PyType> memberReturnTypes = new ArrayList<>();
final PyUnionType unionType = (PyUnionType)type;
for (PyType member : unionType.getMembers()) {
memberReturnTypes.add(getGeneratorReturnType(member, context));
memberReturnTypes.add(Ref.deref(getGeneratorReturnType(member, context)));
}
return PyUnionType.union(memberReturnTypes);
return Ref.create(PyUnionType.union(memberReturnTypes));
}
return null;
}
@@ -0,0 +1,5 @@
from typing import Any
async def mycoroutine() -> Any:
pass
@@ -1,18 +1,4 @@
/*
* Copyright 2000-2017 JetBrains s.r.o.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
// Copyright 2000-2017 JetBrains s.r.o. Use of this source code is governed by the Apache 2.0 license that can be found in the LICENSE file.
package com.jetbrains.python;
import com.intellij.openapi.project.Project;
@@ -22,6 +8,7 @@ import com.jetbrains.python.fixtures.PyTestCase;
import com.jetbrains.python.psi.LanguageLevel;
import com.jetbrains.python.psi.PyExpression;
import com.jetbrains.python.psi.types.TypeEvalContext;
import org.jetbrains.annotations.NotNull;
/**
* @author vlan
@@ -620,6 +607,18 @@ public class Py3TypeTest extends PyTestCase {
);
}
// PY-26847
public void testAwaitOnImportedCoroutine() {
runWithLanguageLevel(
LanguageLevel.PYTHON35,
() -> doMultiFileTest("Any",
"from mycoroutines import mycoroutine\n" +
"\n" +
"async def main():\n" +
" expr = await mycoroutine()")
);
}
private void doTest(final String expectedType, final String text) {
myFixture.configureByText(PythonFileType.INSTANCE, text);
final PyExpression expr = myFixture.findElementByText("expr", PyExpression.class);
@@ -628,4 +627,9 @@ public class Py3TypeTest extends PyTestCase {
assertType(expectedType, expr, TypeEvalContext.codeAnalysis(project, containingFile));
assertType(expectedType, expr, TypeEvalContext.userInitiated(project, containingFile));
}
private void doMultiFileTest(@NotNull String expectedType, @NotNull String text) {
myFixture.copyDirectoryToProject(TEST_DIRECTORY + getTestName(false), "");
doTest(expectedType, text);
}
}