diff --git a/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java b/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java index a46d88403f3a..9783b78b6ea1 100644 --- a/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java +++ b/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java @@ -1268,6 +1268,22 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { 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/PyTypeCheckerInspection.java b/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java index 88ded186db58..ac98642c9709 100644 --- a/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java +++ b/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java @@ -198,7 +198,7 @@ public class PyTypeCheckerInspection extends PyInspection { final PyCallableParameter parameter = entry.getValue(); final PyType expected = parameter.getArgumentType(myTypeEvalContext); final PyType actual = myTypeEvalContext.getType(argument); - final boolean matched = PyTypeChecker.match(expected, actual, myTypeEvalContext, substitutions); + final boolean matched = matchParameterAndArgument(expected, actual, substitutions); result.add(new AnalyzeArgumentResult(argument, expected, substituteGenerics(expected, substitutions), actual, matched)); } final PyCallableParameter positionalContainer = getMappedPositionalContainer(mappedParameters); @@ -220,7 +220,7 @@ public class PyTypeCheckerInspection extends PyInspection { // For an expected type with generics we have to match all the actual types against it in order to do proper generic unification if (PyTypeChecker.hasGenerics(expected, myTypeEvalContext)) { final PyType actual = PyUnionType.union(arguments.stream().map(e -> myTypeEvalContext.getType(e)).collect(Collectors.toList())); - final boolean matched = PyTypeChecker.match(expected, actual, myTypeEvalContext, substitutions); + final boolean matched = matchParameterAndArgument(expected, actual, substitutions); return arguments.stream() .map(argument -> new AnalyzeArgumentResult(argument, expected, expectedWithSubstitutions, actual, matched)) .collect(Collectors.toList()); @@ -229,13 +229,20 @@ public class PyTypeCheckerInspection extends PyInspection { return arguments.stream() .map(argument -> { final PyType actual = myTypeEvalContext.getType(argument); - final boolean matched = PyTypeChecker.match(expected, actual, myTypeEvalContext, substitutions); + final boolean matched = matchParameterAndArgument(expected, actual, substitutions); return new AnalyzeArgumentResult(argument, expected, expectedWithSubstitutions, actual, matched); }) .collect(Collectors.toList()); } } + private boolean matchParameterAndArgument(@Nullable PyType parameterType, + @Nullable PyType argumentType, + @NotNull Map substitutions) { + return PyTypeChecker.match(parameterType, argumentType, myTypeEvalContext, substitutions) && + !PyTypingTypeProvider.matchingProtocolDefinitions(parameterType, argumentType, myTypeEvalContext); + } + @Nullable private PyType substituteGenerics(@Nullable PyType expectedArgumentType, @NotNull Map substitutions) { return PyTypeChecker.hasGenerics(expectedArgumentType, myTypeEvalContext) diff --git a/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspectionProblemRegistrar.java b/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspectionProblemRegistrar.java index eff591d759ab..c501523a37f3 100644 --- a/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspectionProblemRegistrar.java +++ b/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspectionProblemRegistrar.java @@ -22,6 +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.documentation.PythonDocumentationProvider; import com.jetbrains.python.psi.*; import com.jetbrains.python.psi.types.PyClassLikeType; @@ -111,6 +112,10 @@ class PyTypeCheckerInspectionProblemRegistrar { argumentResult.getExpectedTypeAfterSubstitution(), context); + if (PyTypingTypeProvider.matchingProtocolDefinitions(expectedType, actualType, context)) { + return "Only concrete class can be given where " + expectedTypeRepresentation + " is expected"; + } + return String.format("Expected type %s, got '%s' instead", expectedTypeRepresentation, actualTypeName); } diff --git a/python/testData/inspections/PyTypeCheckerInspection/AgainstTypingProtocolDefinition.py b/python/testData/inspections/PyTypeCheckerInspection/AgainstTypingProtocolDefinition.py new file mode 100644 index 000000000000..87a47c26ee60 --- /dev/null +++ b/python/testData/inspections/PyTypeCheckerInspection/AgainstTypingProtocolDefinition.py @@ -0,0 +1,55 @@ +from typing import Protocol, Type + + +class Proto(Protocol): + def proto(self, i: int) -> None: + pass + + +class Concrete1: + def proto(self, i: int) -> None: + pass + + +class Concrete2(Proto): + def proto(self, i: int) -> None: + pass + + +class Concrete3: + def proto(self, i: str) -> None: + pass + + +class Concrete4(Proto): + def proto(self, i: str) -> None: + pass + + +class NewProto(Proto, Protocol): + def new_proto(self, s: str) -> None: + pass + + +def foo(cls: Type[Proto]) -> None: + pass + + +def bar(*classes: Type[Proto]) -> None: + pass + + +foo(Proto) +foo(Concrete1) +foo(Concrete2) +foo(Concrete3) +foo(Concrete4) # matched as inheritor +foo(NewProto) + + +bar(Proto) +bar(Concrete1) +bar(Concrete2) +bar(Concrete3) +bar(Concrete4) # matched as inheritor +bar(NewProto) diff --git a/python/testSrc/com/jetbrains/python/inspections/PyTypeCheckerInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/PyTypeCheckerInspectionTest.java index 7340798045da..9b367326165c 100644 --- a/python/testSrc/com/jetbrains/python/inspections/PyTypeCheckerInspectionTest.java +++ b/python/testSrc/com/jetbrains/python/inspections/PyTypeCheckerInspectionTest.java @@ -510,4 +510,9 @@ public class PyTypeCheckerInspectionTest extends PyInspectionTestCase { public void testTypingProtocolAgainstProtocol() { runWithLanguageLevel(LanguageLevel.PYTHON36, this::doTest); } + + // PY-26628 + public void testAgainstTypingProtocolDefinition() { + runWithLanguageLevel(LanguageLevel.PYTHON35, this::doTest); + } }