diff --git a/python/src/com/jetbrains/python/codeInsight/typing/PyProtocols.kt b/python/src/com/jetbrains/python/codeInsight/typing/PyProtocols.kt index 7bec4f27032b..3f2527dedf9e 100644 --- a/python/src/com/jetbrains/python/codeInsight/typing/PyProtocols.kt +++ b/python/src/com/jetbrains/python/codeInsight/typing/PyProtocols.kt @@ -4,6 +4,7 @@ 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.PyClass import com.jetbrains.python.psi.PyTypedElement import com.jetbrains.python.psi.resolve.PyResolveContext import com.jetbrains.python.psi.resolve.RatedResolveResult @@ -17,6 +18,10 @@ fun isProtocol(classLikeType: PyClassLikeType, context: TypeEvalContext): Boolea return classLikeType.getSuperClassTypes(context).any { type -> PROTOCOL == type?.classQName } } +fun isProtocol(cls: PyClass, context: TypeEvalContext): Boolean { + return cls.getSuperClassTypes(context).any { type -> PROTOCOL == type?.classQName } +} + fun matchingProtocolDefinitions(expected: PyType?, actual: PyType?, context: TypeEvalContext) = expected is PyClassLikeType && actual is PyClassLikeType && expected.isDefinition && diff --git a/python/src/com/jetbrains/python/inspections/PyAbstractClassInspection.java b/python/src/com/jetbrains/python/inspections/PyAbstractClassInspection.java index dcb12155b567..a7533f04cd0e 100644 --- a/python/src/com/jetbrains/python/inspections/PyAbstractClassInspection.java +++ b/python/src/com/jetbrains/python/inspections/PyAbstractClassInspection.java @@ -8,6 +8,8 @@ import com.intellij.psi.PsiElementVisitor; import com.jetbrains.python.PyBundle; import com.jetbrains.python.PyNames; import com.jetbrains.python.codeInsight.override.PyOverrideImplementUtil; +import com.jetbrains.python.codeInsight.typing.PyProtocolsKt; +import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider; import com.jetbrains.python.inspections.quickfix.PyImplementMethodsQuickFix; import com.jetbrains.python.psi.*; import com.jetbrains.python.psi.types.PyClassLikeType; @@ -43,7 +45,7 @@ public class PyAbstractClassInspection extends PyInspection { @Override public void visitPyClass(PyClass pyClass) { - if (isAbstract(pyClass)) { + if (isAbstract(pyClass) || PyProtocolsKt.isProtocol(pyClass, myTypeEvalContext)) { return; } final List toImplement = PyOverrideImplementUtil.getAllSuperAbstractMethods(pyClass, myTypeEvalContext); diff --git a/python/testData/inspections/PyAbstractClassInspection/typingProtocolSubclass.py b/python/testData/inspections/PyAbstractClassInspection/typingProtocolSubclass.py new file mode 100644 index 000000000000..94bef853f8c3 --- /dev/null +++ b/python/testData/inspections/PyAbstractClassInspection/typingProtocolSubclass.py @@ -0,0 +1,118 @@ +from typing import Protocol +from abc import abstractmethod + + +# ------------------ +class MP1(Protocol): + name: str + + +class CMP1(MP1): # ok + pass + + +def test_mp1(mp1: MP1): + pass + + +test_mp1(CMP1()) + + +# ------------------ +class MP2(Protocol): + name: str = "name" + + +class CMP2(MP2): # ok + pass + + +def test_mp2(mp2: MP2): + pass + + +test_mp2(CMP2()) + + +# ------------------ +class MP3(Protocol): + def foo(self) -> int: + return 42 + + +class CMP3(MP3): # ok + pass + + +def test_mp3(mp3: MP3): + pass + + +test_mp3(CMP3()) + + +# ------------------ +class MP4(Protocol): + def foo(self) -> int: + pass + + +class CMP4(MP4): # ok + pass + + +def test_mp4(mp4: MP4): + pass + + +test_mp4(CMP4()) + + +# ------------------ +class MP5(Protocol): + @abstractmethod + def foo(self) -> int: + raise NotImplementedError + + +class CMP5(MP5): # fail + pass + + +def test_mp5(mp5: MP5): + pass + + +test_mp5(CMP5()) + + +# ------------------ +class MP6(Protocol): + @abstractmethod + def foo(self) -> int: + raise NotImplementedError + + +class CMP6(MP6): # ok + @abstractmethod + def bar(self) -> int: + raise NotImplementedError + + +def test_mp6(mp6: MP6): + pass + + +test_mp6(CMP6()) + + +# ------------------ +class MP7(Protocol): + @abstractmethod + def foo(self) -> int: + raise NotImplementedError + + +class PMP7(MP7, Protocol): # ok + def bar(self) -> int: + pass \ No newline at end of file diff --git a/python/testSrc/com/jetbrains/python/inspections/PyAbstractClassInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/PyAbstractClassInspectionTest.java index b6d3e6ce860d..46aad9a4b54f 100644 --- a/python/testSrc/com/jetbrains/python/inspections/PyAbstractClassInspectionTest.java +++ b/python/testSrc/com/jetbrains/python/inspections/PyAbstractClassInspectionTest.java @@ -65,6 +65,11 @@ public class PyAbstractClassInspectionTest extends PyInspectionTestCase { doTest(); } + // PY-26628 + public void testTypingProtocolSubclass() { + runWithLanguageLevel(LanguageLevel.PYTHON37, this::doTest); + } + @NotNull @Override protected Class getInspectionClass() {