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 extends RatedResolveResult> 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 extends RatedResolveResult> 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 extends PyInspection> getInspectionClass() {
+ return PyProtocolInspection.class;
+ }
+}