From 68a7b6040ee804595e89018f091ad217f39e5b35 Mon Sep 17 00:00:00 2001 From: Semyon Proshev Date: Fri, 26 Jan 2018 20:50:44 +0300 Subject: [PATCH] Inspect protocol subclass types (PY-26628) --- .../PyProtocolInspection.html | 5 ++ python/src/META-INF/python-core-common.xml | 3 +- .../python/codeInsight/typing/PyProtocols.kt | 54 +++++++++++++++ .../typing/PyTypingTypeProvider.java | 21 ------ .../inspections/PyProtocolInspection.kt | 67 +++++++++++++++++++ .../inspections/PyTypeCheckerInspection.java | 3 +- ...TypeCheckerInspectionProblemRegistrar.java | 6 +- .../python/psi/types/PyTypeChecker.java | 60 ++++++++--------- .../incompatibleProtocolSubclass.py | 60 +++++++++++++++++ .../validProtocolSubclass.py | 31 +++++++++ .../AgainstTypingProtocolDefinition.py | 8 +-- .../inspections/PyProtocolInspectionTest.java | 38 +++++++++++ 12 files changed, 294 insertions(+), 62 deletions(-) create mode 100644 python/resources/inspectionDescriptions/PyProtocolInspection.html create mode 100644 python/src/com/jetbrains/python/codeInsight/typing/PyProtocols.kt create mode 100644 python/src/com/jetbrains/python/inspections/PyProtocolInspection.kt create mode 100644 python/testData/inspections/PyProtocolInspection/incompatibleProtocolSubclass.py create mode 100644 python/testData/inspections/PyProtocolInspection/validProtocolSubclass.py create mode 100644 python/testSrc/com/jetbrains/python/inspections/PyProtocolInspectionTest.java diff --git a/python/resources/inspectionDescriptions/PyProtocolInspection.html b/python/resources/inspectionDescriptions/PyProtocolInspection.html new file mode 100644 index 000000000000..2200d3d9ccb3 --- /dev/null +++ b/python/resources/inspectionDescriptions/PyProtocolInspection.html @@ -0,0 +1,5 @@ + + +This inspection detects invalid definitions and usages of protocols introduced in PEP-544. + + \ No newline at end of file diff --git a/python/src/META-INF/python-core-common.xml b/python/src/META-INF/python-core-common.xml index 4a5e192e4c23..3fde43ea74f9 100644 --- a/python/src/META-INF/python-core-common.xml +++ b/python/src/META-INF/python-core-common.xml @@ -417,7 +417,8 @@ - + + diff --git a/python/src/com/jetbrains/python/codeInsight/typing/PyProtocols.kt b/python/src/com/jetbrains/python/codeInsight/typing/PyProtocols.kt new file mode 100644 index 000000000000..7bec4f27032b --- /dev/null +++ b/python/src/com/jetbrains/python/codeInsight/typing/PyProtocols.kt @@ -0,0 +1,54 @@ +// Copyright 2000-2018 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.codeInsight.typing + +import com.intellij.util.containers.isNullOrEmpty +import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider.PROTOCOL +import com.jetbrains.python.psi.AccessDirection +import com.jetbrains.python.psi.PyTypedElement +import com.jetbrains.python.psi.resolve.PyResolveContext +import com.jetbrains.python.psi.resolve.RatedResolveResult +import com.jetbrains.python.psi.types.PyClassLikeType +import com.jetbrains.python.psi.types.PyClassType +import com.jetbrains.python.psi.types.PyType +import com.jetbrains.python.psi.types.TypeEvalContext + + +fun isProtocol(classLikeType: PyClassLikeType, context: TypeEvalContext): Boolean { + return classLikeType.getSuperClassTypes(context).any { type -> PROTOCOL == type?.classQName } +} + +fun matchingProtocolDefinitions(expected: PyType?, actual: PyType?, context: TypeEvalContext) = expected is PyClassLikeType && + actual is PyClassLikeType && + expected.isDefinition && + actual.isDefinition && + isProtocol(expected, context) && + isProtocol(actual, context) + +fun inspectProtocolSubclass(protocol: PyClassType, + subclass: PyClassType, + context: TypeEvalContext, + callback: InspectingProtocolSubclassCallback) { + val subclassAsInstance = subclass.toInstance() + val resolveContext = PyResolveContext.noImplicits().withTypeEvalContext(context) + var result = true + + protocol.toInstance().visitMembers( + { e -> + if (result && e is PyTypedElement) { + val name = e.name ?: return@visitMembers result + val resolveResults = subclassAsInstance.resolveMember(name, null, AccessDirection.READ, resolveContext) + + result = if (resolveResults.isNullOrEmpty()) callback.onUnresolved(e) else callback.onResolved(e, resolveResults!!) + } + + result + }, + true, + context + ) +} + +interface InspectingProtocolSubclassCallback { + fun onUnresolved(protocolElement: PyTypedElement): Boolean + fun onResolved(protocolElement: PyTypedElement, subclassElements: List): Boolean +} \ No newline at end of file diff --git a/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java b/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java index 9783b78b6ea1..6d146d9b63b3 100644 --- a/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java +++ b/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java @@ -1263,27 +1263,6 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { return false; } - public static boolean isProtocol(@NotNull PyClassLikeType classLikeType, @NotNull TypeEvalContext context) { - return ContainerUtil.exists(classLikeType.getSuperClassTypes(context), - superClass -> superClass != null && PROTOCOL.equals(superClass.getClassQName())); - } - - public static boolean matchingProtocolDefinitions(@Nullable PyType expected, @Nullable PyType actual, @NotNull TypeEvalContext context) { - if (expected instanceof PyClassLikeType && actual instanceof PyClassLikeType) { - final PyClassLikeType expectedClassLikeType = (PyClassLikeType)expected; - final PyClassLikeType actualClassLikeType = (PyClassLikeType)actual; - - if (expectedClassLikeType.isDefinition() && - actualClassLikeType.isDefinition() && - isProtocol(expectedClassLikeType, context) && - isProtocol(actualClassLikeType, context)) { - return true; - } - } - - return false; - } - @NotNull private static String getOpenMode(@NotNull PyFunction function, @NotNull PyCallExpression call, @NotNull TypeEvalContext context) { final Map arguments = diff --git a/python/src/com/jetbrains/python/inspections/PyProtocolInspection.kt b/python/src/com/jetbrains/python/inspections/PyProtocolInspection.kt new file mode 100644 index 000000000000..c96b14568965 --- /dev/null +++ b/python/src/com/jetbrains/python/inspections/PyProtocolInspection.kt @@ -0,0 +1,67 @@ +// Copyright 2000-2018 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.inspections + +import com.intellij.codeInspection.LocalInspectionToolSession +import com.intellij.codeInspection.ProblemsHolder +import com.intellij.psi.PsiElementVisitor +import com.intellij.psi.PsiNameIdentifierOwner +import com.jetbrains.python.codeInsight.typing.InspectingProtocolSubclassCallback +import com.jetbrains.python.codeInsight.typing.inspectProtocolSubclass +import com.jetbrains.python.codeInsight.typing.isProtocol +import com.jetbrains.python.psi.PyClass +import com.jetbrains.python.psi.PyTypedElement +import com.jetbrains.python.psi.resolve.RatedResolveResult +import com.jetbrains.python.psi.types.PyClassType +import com.jetbrains.python.psi.types.PyTypeChecker + +class PyProtocolInspection : PyInspection() { + + override fun buildVisitor(holder: ProblemsHolder, + isOnTheFly: Boolean, + session: LocalInspectionToolSession): PsiElementVisitor = Visitor(holder, session) + + private class Visitor(holder: ProblemsHolder, session: LocalInspectionToolSession) : PyInspectionVisitor(holder, session) { + + override fun visitPyClass(node: PyClass?) { + super.visitPyClass(node) + + val type = node?.let { myTypeEvalContext.getType(it) } + if (type is PyClassType) { + type + .getSuperClassTypes(myTypeEvalContext) + .asSequence() + .filterIsInstance() + .filter { isProtocol(it, myTypeEvalContext) } + .forEach { protocol -> + inspectProtocolSubclass( + protocol, + type, + myTypeEvalContext, + object : InspectingProtocolSubclassCallback { + override fun onUnresolved(protocolElement: PyTypedElement): Boolean { + return true + } + + override fun onResolved(protocolElement: PyTypedElement, subclassElements: List): Boolean { + val expectedMemberType = myTypeEvalContext.getType(protocolElement) + + subclassElements + .asSequence() + .map { it.element } + .filterIsInstance() + .filter { it.containingFile == node.containingFile } + .filterNot { PyTypeChecker.match(expectedMemberType, myTypeEvalContext.getType(it), myTypeEvalContext) } + .forEach { + val place = if (it is PsiNameIdentifierOwner) it.nameIdentifier else it + registerProblem(place, "Type of '${it.name}' is incompatible with '${protocol.name}'") + } + + return true + } + } + ) + } + } + } + } +} \ No newline at end of file diff --git a/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java b/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java index ac98642c9709..8c5286c2b157 100644 --- a/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java +++ b/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java @@ -10,6 +10,7 @@ import com.intellij.psi.PsiElementVisitor; import com.jetbrains.python.PyNames; import com.jetbrains.python.codeInsight.controlflow.ScopeOwner; import com.jetbrains.python.codeInsight.dataflow.scope.ScopeUtil; +import com.jetbrains.python.codeInsight.typing.PyProtocolsKt; import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider; import com.jetbrains.python.documentation.PythonDocumentationProvider; import com.jetbrains.python.inspections.quickfix.PyMakeFunctionReturnTypeQuickFix; @@ -240,7 +241,7 @@ public class PyTypeCheckerInspection extends PyInspection { @Nullable PyType argumentType, @NotNull Map substitutions) { return PyTypeChecker.match(parameterType, argumentType, myTypeEvalContext, substitutions) && - !PyTypingTypeProvider.matchingProtocolDefinitions(parameterType, argumentType, myTypeEvalContext); + !PyProtocolsKt.matchingProtocolDefinitions(parameterType, argumentType, myTypeEvalContext); } @Nullable diff --git a/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspectionProblemRegistrar.java b/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspectionProblemRegistrar.java index c501523a37f3..90fea3384212 100644 --- a/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspectionProblemRegistrar.java +++ b/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspectionProblemRegistrar.java @@ -22,7 +22,7 @@ import com.intellij.psi.PsiElement; import com.intellij.util.ObjectUtils; import com.intellij.util.containers.ContainerUtil; import com.intellij.xml.util.XmlStringUtil; -import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider; +import com.jetbrains.python.codeInsight.typing.PyProtocolsKt; import com.jetbrains.python.documentation.PythonDocumentationProvider; import com.jetbrains.python.psi.*; import com.jetbrains.python.psi.types.PyClassLikeType; @@ -112,8 +112,8 @@ class PyTypeCheckerInspectionProblemRegistrar { argumentResult.getExpectedTypeAfterSubstitution(), context); - if (PyTypingTypeProvider.matchingProtocolDefinitions(expectedType, actualType, context)) { - return "Only concrete class can be given where " + expectedTypeRepresentation + " is expected"; + if (PyProtocolsKt.matchingProtocolDefinitions(expectedType, actualType, context)) { + return "Only concrete class can be used where " + expectedTypeRepresentation + " protocol is expected"; } return String.format("Expected type %s, got '%s' instead", expectedTypeRepresentation, actualTypeName); diff --git a/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java b/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java index 23ab211905d4..7561f8c91c40 100644 --- a/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java +++ b/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java @@ -11,7 +11,8 @@ import com.intellij.util.containers.ContainerUtil; import com.jetbrains.python.PyNames; import com.jetbrains.python.codeInsight.dataflow.scope.ScopeUtil; import com.jetbrains.python.codeInsight.stdlib.PyNamedTupleType; -import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider; +import com.jetbrains.python.codeInsight.typing.InspectingProtocolSubclassCallback; +import com.jetbrains.python.codeInsight.typing.PyProtocolsKt; import com.jetbrains.python.psi.*; import com.jetbrains.python.psi.impl.PyBuiltinCache; import com.jetbrains.python.psi.impl.PyTypeProvider; @@ -200,46 +201,41 @@ public class PyTypeChecker { return match(superTupleType.getIteratedItemType(), subTupleType.getIteratedItemType(), context, substitutions, recursive, matching); } } - else if (PyTypingTypeProvider.isProtocol(expectedClassType, context) && !matchClasses(superClass, subClass, context)) { + else if (PyProtocolsKt.isProtocol(expectedClassType, context) && !matchClasses(superClass, subClass, context)) { if (expected instanceof PyCollectionType && !matchGenerics((PyCollectionType)expected, actual, context, substitutions, recursive, matching)) { return false; } final boolean[] result = new boolean[]{true}; - final PyClassLikeType actualAsInstance = actualClassType.toInstance(); - final PyResolveContext resolveContext = PyResolveContext.noImplicits().withTypeEvalContext(context); - expectedClassType.toInstance().visitMembers( - e -> { - if (result[0]) { - final PyTypedElement element = as(e, PyTypedElement.class); - final String name = element == null ? null : element.getName(); - - if (element != null && name != null) { - final List resolveResults = - actualAsInstance.resolveMember(name, null, AccessDirection.READ, resolveContext); - - if (ContainerUtil.isEmpty(resolveResults)) { - result[0] = false; - } - else { - final PyType expectedMemberType = context.getType(element); - - result[0] = StreamEx - .of(resolveResults) - .map(ResolveResult::getElement) - .select(PyTypedElement.class) - .map(context::getType) - .anyMatch(actualMemberType -> match(expectedMemberType, actualMemberType, context, substitutions, recursive, matching)); - } - } + PyProtocolsKt.inspectProtocolSubclass( + expectedClassType, + actualClassType, + context, + new InspectingProtocolSubclassCallback() { + @Override + public boolean onUnresolved(@NotNull PyTypedElement protocolElement) { + result[0] = false; + return false; } - return result[0]; - }, - true, - context + @Override + public boolean onResolved(@NotNull PyTypedElement protocolElement, @NotNull List subclassElements) { + final PyType protocolElementType = context.getType(protocolElement); + + result[0] = StreamEx + .of(subclassElements) + .map(ResolveResult::getElement) + .select(PyTypedElement.class) + .map(context::getType) + .anyMatch( + subclassElementType -> match(protocolElementType, subclassElementType, context, substitutions, recursive, matching) + ); + + return result[0]; + } + } ); return result[0]; diff --git a/python/testData/inspections/PyProtocolInspection/incompatibleProtocolSubclass.py b/python/testData/inspections/PyProtocolInspection/incompatibleProtocolSubclass.py new file mode 100644 index 000000000000..a8fc8b7508a5 --- /dev/null +++ b/python/testData/inspections/PyProtocolInspection/incompatibleProtocolSubclass.py @@ -0,0 +1,60 @@ +from typing import Protocol + + +class MyProtocol(Protocol): + attr: int + + def func(self, p: int) -> str: + pass + + +class MyClass1(MyProtocol): + def __init__(self, attr: int) -> None: + self.attr = attr + + def func(self, p: str) -> int: + pass + + +class MyClass2(MyProtocol): + def __init__(self, attr: str) -> None: + self.attr = attr # mypy says nothing + + def func(self, p: int) -> str: + pass + + +class MyClass3(MyProtocol): + def __init__(self, attr: str) -> None: + self.attr = attr # mypy says nothing + + def func(self, p: str) -> int: + pass + + +class MyClass4(MyProtocol): + attr: int + + def func(self, p: str) -> int: + pass + + +class MyClass5(MyProtocol): + attr: str + + def func(self, p: int) -> str: + pass + + +class MyClass6(MyProtocol): + attr: str + + def func(self, p: str) -> int: + pass + + +class HisProtocol(MyProtocol, Protocol): + attr: str + + def func(self, p: str) -> int: + pass diff --git a/python/testData/inspections/PyProtocolInspection/validProtocolSubclass.py b/python/testData/inspections/PyProtocolInspection/validProtocolSubclass.py new file mode 100644 index 000000000000..a04e5674707a --- /dev/null +++ b/python/testData/inspections/PyProtocolInspection/validProtocolSubclass.py @@ -0,0 +1,31 @@ +from typing import Protocol + + +class MyProtocol(Protocol): + attr: int + + def func(self, p: int) -> str: + pass + + +class MyClass1(MyProtocol): + def __init__(self, attr: int) -> None: + self.attr = attr + + + def func(self, p: int) -> str: + pass + + +class MyClass2(MyProtocol): + attr: int + + def func(self, p: int) -> str: + pass + + +class HisProtocol(MyProtocol, Protocol): + attr: int + + def func(self, p: int) -> str: + pass \ No newline at end of file diff --git a/python/testData/inspections/PyTypeCheckerInspection/AgainstTypingProtocolDefinition.py b/python/testData/inspections/PyTypeCheckerInspection/AgainstTypingProtocolDefinition.py index 87a47c26ee60..532139bd6684 100644 --- a/python/testData/inspections/PyTypeCheckerInspection/AgainstTypingProtocolDefinition.py +++ b/python/testData/inspections/PyTypeCheckerInspection/AgainstTypingProtocolDefinition.py @@ -39,17 +39,17 @@ def bar(*classes: Type[Proto]) -> None: pass -foo(Proto) +foo(Proto) foo(Concrete1) foo(Concrete2) foo(Concrete3) foo(Concrete4) # matched as inheritor -foo(NewProto) +foo(NewProto) -bar(Proto) +bar(Proto) bar(Concrete1) bar(Concrete2) bar(Concrete3) bar(Concrete4) # matched as inheritor -bar(NewProto) +bar(NewProto) diff --git a/python/testSrc/com/jetbrains/python/inspections/PyProtocolInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/PyProtocolInspectionTest.java new file mode 100644 index 000000000000..8983e454af75 --- /dev/null +++ b/python/testSrc/com/jetbrains/python/inspections/PyProtocolInspectionTest.java @@ -0,0 +1,38 @@ +// Copyright 2000-2018 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.inspections; + +import com.intellij.testFramework.LightProjectDescriptor; +import com.jetbrains.python.fixtures.PyInspectionTestCase; +import com.jetbrains.python.psi.LanguageLevel; +import org.jetbrains.annotations.NotNull; +import org.jetbrains.annotations.Nullable; + +public class PyProtocolInspectionTest extends PyInspectionTestCase { + + // PY-26628 + public void testValidProtocolSubclass() { + doTest(); + } + + // PY-26628 + public void testIncompatibleProtocolSubclass() { + doTest(); + } + + @Override + protected void doTest() { + runWithLanguageLevel(LanguageLevel.PYTHON37, () -> super.doTest()); + } + + @Nullable + @Override + protected LightProjectDescriptor getProjectDescriptor() { + return ourPy3Descriptor; + } + + @NotNull + @Override + protected Class getInspectionClass() { + return PyProtocolInspection.class; + } +}