PY-24273, PY-53703 Support for functions annotated with typing.NoReturn and typing.Never

Functions annotated with `NoReturn` and `Never` now taken into account in the Control Flow Graph building process, and the code after calling such functions is treated as unreachable.

Merge-request: IJ-MR-105973
Merged-by: Daniil Kalinin <Daniil.Kalinin@jetbrains.com>

GitOrigin-RevId: ef5840ae6e593498fc334dc9bd2daadccebf2b13
This commit is contained in:
Daniil Kalinin
2023-06-13 22:08:30 +00:00
committed by intellij-monorepo-bot
parent 7f3964f4c7
commit 45bb1fffb8
27 changed files with 177 additions and 125 deletions
@@ -29,11 +29,12 @@ import com.intellij.psi.util.QualifiedName;
import com.intellij.util.containers.ContainerUtil;
import com.jetbrains.python.PyNames;
import com.jetbrains.python.PyTokenTypes;
import com.jetbrains.python.codeInsight.dataflow.scope.ScopeUtil;
import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider;
import com.jetbrains.python.psi.*;
import com.jetbrains.python.psi.impl.ParamHelper;
import com.jetbrains.python.psi.impl.PyAugAssignmentStatementNavigator;
import com.jetbrains.python.psi.impl.PyEvaluator;
import com.jetbrains.python.psi.impl.PyImportStatementNavigator;
import com.jetbrains.python.psi.impl.*;
import com.jetbrains.python.psi.resolve.PyResolveUtil;
import com.jetbrains.python.psi.types.TypeEvalContext;
import kotlin.Triple;
import one.util.streamex.StreamEx;
import org.jetbrains.annotations.NotNull;
@@ -46,8 +47,6 @@ public class PyControlFlowBuilder extends PyRecursiveElementVisitor {
@NotNull
private static final Set<String> EXCEPTION_SUPPRESSORS = ImmutableSet.of("suppress", "assertRaises", "assertRaisesRegex");
private static final Set<String> KNOWN_NORETURNS = ImmutableSet.of("sys.exit", "exit", "pytest.fail");
private final ControlFlowBuilder myBuilder = new ControlFlowBuilder();
public ControlFlow buildControlFlow(@NotNull final ScopeOwner owner) {
@@ -135,7 +134,7 @@ public class PyControlFlowBuilder extends PyRecursiveElementVisitor {
public void visitPyCallExpression(final @NotNull PyCallExpression node) {
final PyExpression callee = node.getCallee();
// Flow abrupted
if (callee != null && assumeDeadEnd(callee)) {
if (callee != null && isCallOfNoReturnFunction(callee)) {
callee.accept(this);
for (PyExpression expression : node.getArguments()) {
expression.accept(this);
@@ -227,7 +226,7 @@ public class PyControlFlowBuilder extends PyRecursiveElementVisitor {
public void visitPyDelStatement(@NotNull PyDelStatement node) {
myBuilder.startNode(node);
for (PyExpression target : node.getTargets()) {
if (target instanceof PyReferenceExpression expr){
if (target instanceof PyReferenceExpression expr) {
myBuilder.addNode(ReadWriteInstruction.newInstruction(myBuilder, target, expr.getName(), ReadWriteInstruction.ACCESS.DELETE));
PyExpression qualifier = expr.getQualifier();
if (qualifier != null) {
@@ -993,23 +992,29 @@ public class PyControlFlowBuilder extends PyRecursiveElementVisitor {
if (target != null) target.accept(this);
}
private static boolean assumeDeadEnd(final @NotNull PyExpression callee) {
String repr = PyUtil.getReadableRepr(callee, true);
if (KNOWN_NORETURNS.contains(repr)) {
return true;
}
/* Since we can't fully resolve the call during the building of the control flow graph,
* here we make an assumption that the class which contains self.fail() call is the real
* test class and self.fail() is actually unittest.TestCase.fail() call which leads to flow abruption (see PY-23859).
* This approach does not completely eliminate false positives, but it helps to reduce their number. */
if (repr.equals("self.fail")) {
PyClass clazz = PsiTreeUtil.getParentOfType(callee, PyClass.class);
if (clazz != null && clazz.getName() != null) {
String className = clazz.getName();
boolean classNameContainsTest = className.contains("Test");
if (classNameContainsTest) {
private static boolean isCallOfNoReturnFunction(@NotNull PyExpression callee) {
if (callee instanceof PyReferenceExpression expression) {
QualifiedName qName = expression.asQualifiedName();
if (qName == null) {
return false;
}
ScopeOwner scopeOwner = ScopeUtil.getScopeOwner(expression);
// Flow-insensitive context is required to prevent recursive control flow access during the resolve process
TypeEvalContext context = TypeEvalContext.codeInsightFallback(callee.getProject());
while (scopeOwner != null) {
boolean resolvesToNoReturnOrNever = StreamEx
.of(PyResolveUtil.resolveQualifiedNameInScope(qName, scopeOwner, context))
.select(PyFunction.class)
.anyMatch(function -> PyTypingTypeProvider.isNoReturn(function, context));
if (resolvesToNoReturnOrNever) {
return true;
}
scopeOwner = ScopeUtil.getScopeOwner(scopeOwner);
}
}
return false;
@@ -43,7 +43,6 @@ import org.jetbrains.annotations.NotNull;
import org.jetbrains.annotations.Nullable;
import java.util.*;
import java.util.function.BiFunction;
import java.util.function.Function;
import java.util.regex.Matcher;
import java.util.regex.Pattern;
@@ -91,6 +90,9 @@ public class PyTypingTypeProvider extends PyTypeProviderWithCustomContext<PyTypi
public static final String TYPING_EXTENSIONS_CONCATENATE = "typing_extensions.Concatenate";
public static final String OPTIONAL = "typing.Optional";
public static final String NO_RETURN = "typing.NoReturn";
public static final String NEVER = "typing.Never";
public static final String NO_RETURN_EXT = "typing_extensions.NoReturn";
public static final String NEVER_EXT = "typing_extensions.Never";
public static final String FINAL = "typing.Final";
public static final String FINAL_EXT = "typing_extensions.Final";
public static final String LITERAL = "typing.Literal";
@@ -796,7 +798,7 @@ public class PyTypingTypeProvider extends PyTypeProviderWithCustomContext<PyTypi
if (callableType != null) {
return Ref.create(callableType);
}
final Ref<PyType> classVarType = getClassVarType(resolved, context);
final Ref<PyType> classVarType = unwrapTypeModifier(resolved, context, CLASS_VAR);
if (classVarType != null) {
return classVarType;
}
@@ -804,7 +806,7 @@ public class PyTypingTypeProvider extends PyTypeProviderWithCustomContext<PyTypi
if (classObjType != null) {
return classObjType;
}
final Ref<PyType> finalType = getFinalType(resolved, context);
final Ref<PyType> finalType = unwrapTypeModifier(resolved, context, FINAL, FINAL_EXT);
if (finalType != null) {
return finalType;
}
@@ -932,19 +934,6 @@ public class PyTypingTypeProvider extends PyTypeProviderWithCustomContext<PyTypi
return null;
}
@Nullable
private static Ref<PyType> getClassVarType(@NotNull PsiElement resolved, @NotNull Context context) {
if (resolved instanceof PySubscriptionExpression subscriptionExpr) {
if (resolvesToClassVar(subscriptionExpr.getOperand(), context.getTypeContext())) {
final PyExpression indexExpr = subscriptionExpr.getIndexExpression();
if (indexExpr != null) {
return getType(indexExpr, context);
}
}
}
return null;
}
@Nullable
private static Ref<PyType> getAliasedType(@NotNull PsiElement resolved, @NotNull Context context) {
if (resolved instanceof PyReferenceExpression && ((PyReferenceExpression)resolved).asQualifiedName() != null) {
@@ -1119,9 +1108,9 @@ public class PyTypingTypeProvider extends PyTypeProviderWithCustomContext<PyTypi
}
@Nullable
private static Ref<PyType> getFinalType(@NotNull PsiElement resolved, @NotNull Context context) {
private static Ref<PyType> unwrapTypeModifier(@NotNull PsiElement resolved, @NotNull Context context, String... type) {
if (resolved instanceof PySubscriptionExpression subscriptionExpr) {
if (resolvesToFinal(subscriptionExpr.getOperand(), context.getTypeContext())) {
if (resolvesToQualifiedNames(subscriptionExpr.getOperand(), context.getTypeContext(), type)) {
final PyExpression indexExpr = subscriptionExpr.getIndexExpression();
if (indexExpr != null) {
return getType(indexExpr, context);
@@ -1132,24 +1121,24 @@ public class PyTypingTypeProvider extends PyTypeProviderWithCustomContext<PyTypi
return null;
}
private static <T extends PyTypeCommentOwner & PyAnnotationOwner> boolean isSpecialModifierImpl(@NotNull T owner,
@NotNull TypeEvalContext context,
@NotNull BiFunction<PyExpression, TypeEvalContext, Boolean> resolver) {
private static <T extends PyTypeCommentOwner & PyAnnotationOwner> boolean typeHintedWithName(@NotNull T owner,
@NotNull TypeEvalContext context,
String... names) {
final PyExpression annotation = getAnnotationValue(owner, context);
if (annotation instanceof PySubscriptionExpression) {
return resolver.apply(((PySubscriptionExpression)annotation).getOperand(), context);
return resolvesToQualifiedNames(((PySubscriptionExpression)annotation).getOperand(), context, names);
}
else if (annotation instanceof PyReferenceExpression) {
return resolver.apply(annotation, context);
return resolvesToQualifiedNames(annotation, context, names);
}
final String typeCommentValue = owner.getTypeCommentAnnotation();
final PyExpression typeComment = typeCommentValue == null ? null : toExpression(typeCommentValue, owner);
if (typeComment instanceof PySubscriptionExpression) {
return resolver.apply(((PySubscriptionExpression)typeComment).getOperand(), context);
return resolvesToQualifiedNames(((PySubscriptionExpression)typeComment).getOperand(), context, names);
}
else if (typeComment instanceof PyReferenceExpression) {
return resolver.apply(typeComment, context);
return resolvesToQualifiedNames(typeComment, context, names);
}
return false;
@@ -1161,25 +1150,23 @@ public class PyTypingTypeProvider extends PyTypeProviderWithCustomContext<PyTypi
}
public static <T extends PyTypeCommentOwner & PyAnnotationOwner> boolean isFinal(@NotNull T owner, @NotNull TypeEvalContext context) {
return PyUtil.getParameterizedCachedValue(owner, context, p -> isSpecialModifierImpl(owner, p, (e, c) -> {
return resolvesToFinal(e, c);
}));
}
private static boolean resolvesToFinal(@NotNull PyExpression expression, @NotNull TypeEvalContext context) {
final var qualifiedNames = resolveToQualifiedNames(expression, context);
return qualifiedNames.contains(FINAL) || qualifiedNames.contains(FINAL_EXT);
return PyUtil.getParameterizedCachedValue(owner, context, p ->
typeHintedWithName(owner, context, FINAL, FINAL_EXT));
}
public static <T extends PyAnnotationOwner & PyTypeCommentOwner> boolean isClassVar(@NotNull T owner, @NotNull TypeEvalContext context) {
return PyUtil.getParameterizedCachedValue(owner, context, p -> isSpecialModifierImpl(owner, p, (e, c) -> {
return resolvesToClassVar(e, c);
}));
return PyUtil.getParameterizedCachedValue(owner, context, p ->
typeHintedWithName(owner, context, CLASS_VAR));
}
private static boolean resolvesToClassVar(@NotNull PyExpression expression, @NotNull TypeEvalContext context) {
public static boolean isNoReturn(@NotNull PyFunction function, @NotNull TypeEvalContext context) {
return PyUtil.getParameterizedCachedValue(function, context, p ->
typeHintedWithName(function, context, NO_RETURN, NO_RETURN_EXT, NEVER, NEVER_EXT));
}
private static boolean resolvesToQualifiedNames(@NotNull PyExpression expression, @NotNull TypeEvalContext context, String... names) {
final var qualifiedNames = resolveToQualifiedNames(expression, context);
return qualifiedNames.contains(CLASS_VAR);
return ContainerUtil.exists(names, qualifiedNames::contains);
}
@Nullable
@@ -219,7 +219,8 @@ public class PyTypeCheckerInspection extends PyInspection {
}
}
if (PyUtil.isInitMethod(node) && !(getExpectedReturnType(node) instanceof PyNoneType)) {
if (PyUtil.isInitMethod(node) && !(getExpectedReturnType(node) instanceof PyNoneType
|| PyTypingTypeProvider.isNoReturn(node, myTypeEvalContext))) {
registerProblem(annotation != null ? annotation.getValue() : node.getTypeComment(),
PyPsiBundle.message("INSP.type.checker.init.should.return.none"));
}
@@ -644,19 +644,19 @@ public final class PyUtil {
* @param <P> key type
*/
@NotNull
public static <T, P> T getParameterizedCachedValue(@NotNull PsiElement element, @Nullable P param, @NotNull NotNullFunction<P, T> f) {
public static <T, P> T getParameterizedCachedValue(@NotNull PsiElement element, @Nullable P param, @NotNull Function<P, @NotNull T> f) {
final T result = getNullableParameterizedCachedValue(element, param, f);
assert result != null;
return result;
}
/**
* Same as {@link #getParameterizedCachedValue(PsiElement, Object, NotNullFunction)} but allows nulls.
* Same as {@link #getParameterizedCachedValue(PsiElement, Object, Function)} but allows nulls.
*/
@Nullable
public static <T, P> T getNullableParameterizedCachedValue(@NotNull PsiElement element,
@Nullable P param,
@NotNull NullableFunction<P, T> f) {
@NotNull Function<P, @Nullable T> f) {
final CachedValuesManager manager = CachedValuesManager.getManager(element.getProject());
final Map<Optional<P>, Optional<T>> cache = CachedValuesManager.getCachedValue(element, manager.getKeyForClass(f.getClass()), () -> {
// concurrent hash map is a null-hostile collection
@@ -1,6 +0,0 @@
def test_fail():
if True == False:
pytest.fail()
print("should be reported as unreachable")
else:
return 1
@@ -1,11 +0,0 @@
0(1) element: null
1(2) element: PyIfStatement
2(3) READ ACCESS: True
3(4,8) READ ACCESS: False
4(5) element: PyStatementList. Condition: True == False:true
5(6) element: PyExpressionStatement
6(10) READ ACCESS: pytest
7(10) element: PyPrintStatement
8(9) element: PyStatementList. Condition: True == False:false
9(10) element: PyReturnStatement
10() element: null
@@ -1,17 +0,0 @@
0(1) element: null
1(2) element: PyTryExceptStatement
2(3,8) element: PyTryPart
3(4,8) element: PyAssignmentStatement
4(5,8) READ ACCESS: int
5(6,8) element: PySubscriptionExpression
6(7,8) READ ACCESS: sys
7(8,13) WRITE ACCESS: n
8(9) element: PyExceptPart
9(10) READ ACCESS: ValueError
10(11) element: PyPrintStatement
11(12) element: PyExpressionStatement
12(16) READ ACCESS: sys
13(14) element: PyPrintStatement
14(15) READ ACCESS: str
15(16) READ ACCESS: n
16() element: null
@@ -0,0 +1,7 @@
from typing import Never
def stop() -> Never:
raise RuntimeError('no way')
stop()
print("ureachable")
@@ -0,0 +1,10 @@
0(1) element: null
1(2) element: PyFromImportStatement
2(3) WRITE ACCESS: Never
3(4) element: PyFunction('stop')
4(5) READ ACCESS: Never
5(6) WRITE ACCESS: stop
6(7) element: PyExpressionStatement
7(9) READ ACCESS: stop
8(9) element: PyPrintStatement
9() element: null
@@ -0,0 +1,7 @@
from typing import NoReturn
def stop() -> NoReturn:
raise RuntimeError('no way')
stop()
print("ureachable")
@@ -0,0 +1,10 @@
0(1) element: null
1(2) element: PyFromImportStatement
2(3) WRITE ACCESS: NoReturn
3(4) element: PyFunction('stop')
4(5) READ ACCESS: NoReturn
5(6) WRITE ACCESS: stop
6(7) element: PyExpressionStatement
7(9) READ ACCESS: stop
8(9) element: PyPrintStatement
9() element: null
@@ -0,0 +1,19 @@
0(1) element: null
1(2) element: PyImportStatement
2(3) WRITE ACCESS: sys
3(4) element: PyTryExceptStatement
4(5,10) element: PyTryPart
5(6,10) element: PyAssignmentStatement
6(7,10) READ ACCESS: int
7(8,10) element: PySubscriptionExpression
8(9,10) READ ACCESS: sys
9(10,15) WRITE ACCESS: n
10(11) element: PyExceptPart
11(12) READ ACCESS: ValueError
12(13) element: PyPrintStatement
13(14) element: PyExpressionStatement
14(18) READ ACCESS: sys
15(16) element: PyPrintStatement
16(17) READ ACCESS: str
17(18) READ ACCESS: n
18() element: null
@@ -1,8 +0,0 @@
# PY-23859
from unittest import TestCase
class C(TestCase):
def test_1(self):
self.fail()
return -42
@@ -0,0 +1,5 @@
from util import panic
def foo():
panic("Error!")
<warning descr="This code is unreachable">print("Should be reported as unreachable")</warning>
@@ -0,0 +1,5 @@
from typing import NoReturn
def panic(m) -> NoReturn:
print(f'Help: {m}')
raise SystemExit
@@ -0,0 +1,7 @@
from typing import Never
def stop() -> Never:
raise RuntimeError('no way')
stop()
<warning descr="This code is unreachable">print("Should be reported as unreachable")</warning>
@@ -0,0 +1,7 @@
from typing import NoReturn
def stop() -> NoReturn:
raise RuntimeError('no way')
stop()
<warning descr="This code is unreachable">print("Should be reported as unreachable")</warning>
@@ -1,8 +0,0 @@
import pytest
def test_fail():
if True == False:
pytest.fail()
<warning descr="This code is unreachable">print("should be reported as unreachable")</warning>
else:
return 1
@@ -479,17 +479,17 @@ public class PyControlFlowBuilderTest extends LightMarkedTestCase {
}
// PY-7758
public void testControlFlowAbruptedOnExit() {
public void testControlFlowIsAbruptAfterExit() {
doTest();
}
// PY-7758
public void testControlFlowAbruptedOnSysExit() {
public void testControlFlowIsAbruptAfterSysExit() {
doTest();
}
// PY-23859
public void testControlFlowAbruptedOnRealSelfFailAssumedByClassName() {
public void testControlFlowIsAbruptAfterSelfFail() {
final String testName = getTestName(false);
configureByFile(testName + ".py");
final String fullPath = getTestDataPath() + testName + ".txt";
@@ -498,10 +498,17 @@ public class PyControlFlowBuilderTest extends LightMarkedTestCase {
check(fullPath, flow);
}
public void testControlFlowAbruptedOnPytestFail() {
doTestFirstStatement();
// PY-24273
public void testControlFlowIsAbruptAfterNoReturn() {
doTest();
}
// TODO migrate this test class to Python 3 SDK by default to make this test work
// PY-53703
//public void testControlFlowIsAbruptAfterNever() {
// doTest();
//}
private void doTestFirstStatement() {
final String testName = getTestName(false);
configureByFile(testName + ".py");
@@ -1516,4 +1516,19 @@ public class PyTypeCheckerInspectionTest extends PyInspectionTestCase {
group: GroupWithOtherKey[str, int] = {"key": <warning descr="Expected type 'str', got 'int' instead">1</warning>, "group": [], "some_other_key": <warning descr="Expected type 'int', got 'str' instead">''</warning>}""")
);
}
// PY-27551
public void testDunderInitAnnotatedWithNoReturn() {
runWithLanguageLevel(
LanguageLevel.getLatest(),
() -> doTestByText("""
from typing import NoReturn
class Test:
def __init__(self) -> NoReturn:
raise Exception()
""")
);
}
}
@@ -219,17 +219,25 @@ public class PyUnreachableCodeInspectionTest extends PyInspectionTestCase {
}
// PY-23859
public void testUnreachableCodeReportedAfterSelfFailInClassContainingTestInName() {
public void testUnreachableCodeReportedAfterSelfFail() {
doTest();
}
// PY-23859
public void testCodeNotReportedAsUnreachableAfterSelfFailInClassNotContainingTestInName() {
// PY-24273
public void testUnreachableCodeReportedAfterNoReturnFunction() {
doTest();
}
public void testUnreachableCodeReportedAfterPytestFail() {
doTest();
// PY-24273
public void testUnreachableCodeReportedAfterImportedNoReturnFunction() {
doMultiFileTest();
}
// PY-53703
public void testUnreachableCodeReportedAfterNever() {
runWithLanguageLevel(LanguageLevel.getLatest(), () -> {
doTest();
});
}
@NotNull