diff --git a/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java b/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java index 6d146d9b63b3..a46d88403f3a 100644 --- a/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java +++ b/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java @@ -1263,6 +1263,11 @@ 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())); + } + @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/psi/types/PyTypeChecker.java b/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java index 5243f9492ed0..23ab211905d4 100644 --- a/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java +++ b/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java @@ -2,13 +2,16 @@ package com.jetbrains.python.psi.types; import com.intellij.openapi.extensions.Extensions; +import com.intellij.openapi.util.Pair; import com.intellij.psi.PsiElement; import com.intellij.psi.PsiFile; +import com.intellij.psi.ResolveResult; import com.intellij.util.ArrayUtil; 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.psi.*; import com.jetbrains.python.psi.impl.PyBuiltinCache; import com.jetbrains.python.psi.impl.PyTypeProvider; @@ -32,7 +35,7 @@ public class PyTypeChecker { } public static boolean match(@Nullable PyType expected, @Nullable PyType actual, @NotNull TypeEvalContext context) { - return match(expected, actual, context, null, true); + return match(expected, actual, context, null, true, new HashSet<>()); } /** @@ -48,11 +51,28 @@ public class PyTypeChecker { */ public static boolean match(@Nullable PyType expected, @Nullable PyType actual, @NotNull TypeEvalContext context, @Nullable Map substitutions) { - return match(expected, actual, context, substitutions, true); + return match(expected, actual, context, substitutions, true, new HashSet<>()); } private static boolean match(@Nullable PyType expected, @Nullable PyType actual, @NotNull TypeEvalContext context, - @Nullable Map substitutions, boolean recursive) { + @Nullable Map substitutions, boolean recursive, + @NotNull Set> matching) { + final Pair types = Pair.create(expected, actual); + if (matching.contains(types)) return true; + + matching.add(types); + final boolean result = matchImpl(expected, actual, context, substitutions, recursive, matching); + matching.remove(types); + + return result; + } + + private static boolean matchImpl(@Nullable PyType expected, + @Nullable PyType actual, + @NotNull TypeEvalContext context, + @Nullable Map substitutions, + boolean recursive, + @NotNull Set> matching) { // TODO: subscriptable types?, module types?, etc. final PyClassType expectedClassType = as(expected, PyClassType.class); final PyClassType actualClassType = as(actual, PyClassType.class); @@ -86,7 +106,7 @@ public class PyTypeChecker { if (generic.isDefinition() && bound instanceof PyInstantiableType) { bound = ((PyInstantiableType)bound).toClass(); } - if (!match(bound, actual, context, substitutions, recursive)) { + if (!match(bound, actual, context, substitutions, recursive, matching)) { return false; } else if (subst != null) { @@ -94,7 +114,7 @@ public class PyTypeChecker { return true; } else if (recursive) { - return match(subst, actual, context, substitutions, false); + return match(subst, actual, context, substitutions, false, matching); } else { return false; @@ -122,12 +142,12 @@ public class PyTypeChecker { final int elementCount = expectedTupleType.getElementCount(); if (!expectedTupleType.isHomogeneous() && consistsOfSameElementNumberTuples(actualUnionType, elementCount)) { - return substituteExpectedElementsWithUnions(expectedTupleType, elementCount, actualUnionType, context, substitutions, recursive); + return substituteExpectedElementsWithUnions(expectedTupleType, elementCount, actualUnionType, context, substitutions, recursive, matching); } } for (PyType m : actualUnionType.getMembers()) { - if (match(expected, m, context, substitutions, recursive)) { + if (match(expected, m, context, substitutions, recursive, matching)) { return true; } } @@ -139,7 +159,7 @@ public class PyTypeChecker { final StreamEx genericTypes = StreamEx.of(expectedUnionTypeMembers).select(PyGenericType.class); for (PyType t : notGenericTypes.append(genericTypes)) { - if (match(t, actual, context, substitutions, recursive)) { + if (match(t, actual, context, substitutions, recursive, matching)) { return true; } } @@ -157,7 +177,7 @@ public class PyTypeChecker { } else { for (int i = 0; i < superTupleType.getElementCount(); i++) { - if (!match(superTupleType.getElementType(i), subTupleType.getElementType(i), context, substitutions, recursive)) { + if (!match(superTupleType.getElementType(i), subTupleType.getElementType(i), context, substitutions, recursive, matching)) { return false; } } @@ -167,7 +187,7 @@ public class PyTypeChecker { else if (superTupleType.isHomogeneous() && !subTupleType.isHomogeneous()) { final PyType expectedElementType = superTupleType.getIteratedItemType(); for (int i = 0; i < subTupleType.getElementCount(); i++) { - if (!match(expectedElementType, subTupleType.getElementType(i), context, substitutions, recursive)) { + if (!match(expectedElementType, subTupleType.getElementType(i), context, substitutions, recursive, matching)) { return false; } } @@ -177,9 +197,53 @@ public class PyTypeChecker { return false; } else { - return match(superTupleType.getIteratedItemType(), subTupleType.getIteratedItemType(), context, substitutions, recursive); + return match(superTupleType.getIteratedItemType(), subTupleType.getIteratedItemType(), context, substitutions, recursive, matching); } } + else if (PyTypingTypeProvider.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)); + } + } + } + + return result[0]; + }, + true, + context + ); + + return result[0]; + } else if (expected instanceof PyCollectionType && actual instanceof PyTupleType) { if (!matchClasses(superClass, subClass, context)) { return false; @@ -188,31 +252,16 @@ public class PyTypeChecker { final PyType superElementType = ((PyCollectionType)expected).getIteratedItemType(); final PyType subElementType = ((PyTupleType)actual).getIteratedItemType(); - if (!match(superElementType, subElementType, context, substitutions, recursive)) { + if (!match(superElementType, subElementType, context, substitutions, recursive, matching)) { return false; } return true; } else if (expected instanceof PyCollectionType) { - if (!matchClasses(superClass, subClass, context)) { - return false; - } - // TODO: Match generic parameters based on the correspondence between the generic parameters of subClass and its base classes - final List superElementTypes = ((PyCollectionType)expected).getElementTypes(); - final PyCollectionType actualCollectionType = as(actual, PyCollectionType.class); - final List subElementTypes = actualCollectionType != null ? - actualCollectionType.getElementTypes() : - Collections.emptyList(); - for (int i = 0; i < superElementTypes.size(); i++) { - final PyType subElementType = i < subElementTypes.size() ? subElementTypes.get(i) : null; - if (!match(superElementTypes.get(i), subElementType, context, substitutions, recursive)) { - return false; - } - } - return true; + return matchClasses(superClass, subClass, context) && + matchGenerics((PyCollectionType)expected, actual, context, substitutions, recursive, matching); } - else if (matchClasses(superClass, subClass, context)) { return true; } @@ -268,12 +317,12 @@ public class PyTypeChecker { final PyCallableParameter expectedParam = expectedParameters.get(i); final PyCallableParameter actualParam = actualParameters.get(i); // TODO: Check named and star params, not only positional ones - if (!match(expectedParam.getType(context), actualParam.getType(context), context, substitutions, recursive)) { + if (!match(expectedParam.getType(context), actualParam.getType(context), context, substitutions, recursive, matching)) { return false; } } } - if (!match(expectedCallable.getReturnType(context), actualCallable.getReturnType(context), context, substitutions, recursive)) { + if (!match(expectedCallable.getReturnType(context), actualCallable.getReturnType(context), context, substitutions, recursive, matching)) { return false; } return true; @@ -319,7 +368,8 @@ public class PyTypeChecker { @NotNull PyUnionType actual, @NotNull TypeEvalContext context, @Nullable Map substitutions, - boolean recursive) { + boolean recursive, + @NotNull Set> matching) { for (int i = 0; i < elementCount; i++) { final int currentIndex = i; @@ -331,7 +381,7 @@ public class PyTypeChecker { .toList() ); - if (!match(expected.getElementType(i), elementType, context, substitutions, recursive)) { + if (!match(expected.getElementType(i), elementType, context, substitutions, recursive, matching)) { return false; } } @@ -339,6 +389,27 @@ public class PyTypeChecker { return true; } + private static boolean matchGenerics(@NotNull PyCollectionType expected, + @NotNull PyType actual, + @NotNull TypeEvalContext context, + @Nullable Map substitutions, + boolean recursive, + @NotNull Set> matching) { + // TODO: Match generic parameters based on the correspondence between the generic parameters of subClass and its base classes + final List superElementTypes = expected.getElementTypes(); + final PyCollectionType actualCollectionType = as(actual, PyCollectionType.class); + final List subElementTypes = actualCollectionType != null ? + actualCollectionType.getElementTypes() : + Collections.emptyList(); + for (int i = 0; i < superElementTypes.size(); i++) { + final PyType subElementType = i < subElementTypes.size() ? subElementTypes.get(i) : null; + if (!match(superElementTypes.get(i), subElementType, context, substitutions, recursive, matching)) { + return false; + } + } + return true; + } + private static boolean matchNumericTypes(PyType expected, PyType actual) { final String superName = expected.getName(); final String subName = actual.getName(); diff --git a/python/testData/inspections/PyTypeCheckerInspection/AgainstGenericTypingProtocol.py b/python/testData/inspections/PyTypeCheckerInspection/AgainstGenericTypingProtocol.py new file mode 100644 index 000000000000..a580666848b3 --- /dev/null +++ b/python/testData/inspections/PyTypeCheckerInspection/AgainstGenericTypingProtocol.py @@ -0,0 +1,29 @@ +from typing import Generic, Protocol, TypeVar + +T = TypeVar("T") + + +class Box1(Protocol[T]): + attr: T + + +class Box2(Protocol, Generic[T]): + attr: T + + +class BoxImpl(Generic[T]): + def __init__(self, attr: T) -> None: + self.attr = attr + + +def b1(b: Box1[int]): + print(b.attr) + + +def b2(b: Box2[int]): + print(b.attr) + + +b3: BoxImpl[int] +b1(b3) +b2(b3) diff --git a/python/testData/inspections/PyTypeCheckerInspection/AgainstMergedTypingProtocols.py b/python/testData/inspections/PyTypeCheckerInspection/AgainstMergedTypingProtocols.py new file mode 100644 index 000000000000..a9ab48286b4d --- /dev/null +++ b/python/testData/inspections/PyTypeCheckerInspection/AgainstMergedTypingProtocols.py @@ -0,0 +1,29 @@ +from typing import Protocol, Sized + + +class SupportsClose(Protocol): + def close(self) -> None: + pass + + +class SizedAndClosable(Sized, SupportsClose, Protocol): + pass + + +class Resource: + def __len__(self) -> int: + return 0 + + def close(self) -> None: + pass + + +def close(sized_and_closeable: SizedAndClosable) -> None: + print(len(sized_and_closeable)) + sized_and_closeable.close() + + +r = Resource() +close(r) + +close(1) diff --git a/python/testData/inspections/PyTypeCheckerInspection/AgainstRecursiveTypingProtocol.py b/python/testData/inspections/PyTypeCheckerInspection/AgainstRecursiveTypingProtocol.py new file mode 100644 index 000000000000..23e5133381e3 --- /dev/null +++ b/python/testData/inspections/PyTypeCheckerInspection/AgainstRecursiveTypingProtocol.py @@ -0,0 +1,27 @@ +from typing import Generic, Iterable, List, Protocol, TypeVar + +T = TypeVar("T") + + +class Traversable(Protocol): + def leaves(self) -> Iterable['Traversable']: + pass + + +class SimpleTree: + def leaves(self) -> List['SimpleTree']: + pass + + +class Tree(Generic[T]): + def leaves(self) -> List['Tree[T]']: + pass + + +def traverse(t: Traversable): + for l in t.leaves(): + traverse(l) + + +traverse(SimpleTree()) +traverse(Tree()) diff --git a/python/testData/inspections/PyTypeCheckerInspection/AgainstTypingProtocol.py b/python/testData/inspections/PyTypeCheckerInspection/AgainstTypingProtocol.py new file mode 100644 index 000000000000..a7197c592204 --- /dev/null +++ b/python/testData/inspections/PyTypeCheckerInspection/AgainstTypingProtocol.py @@ -0,0 +1,25 @@ +from typing import Protocol + + +class SupportsClose(Protocol): + def close(self) -> None: + pass + + +class Resource: + def close(self) -> None: + pass + + +def close(closeable: SupportsClose) -> None: + closeable.close() + + +f = open("a.txt") +close(f) + +r = Resource() +close(r) +close(Resource) + +close(1) diff --git a/python/testData/inspections/PyTypeCheckerInspection/AgainstTypingProtocolWithImplementedMethod.py b/python/testData/inspections/PyTypeCheckerInspection/AgainstTypingProtocolWithImplementedMethod.py new file mode 100644 index 000000000000..5b51de723c38 --- /dev/null +++ b/python/testData/inspections/PyTypeCheckerInspection/AgainstTypingProtocolWithImplementedMethod.py @@ -0,0 +1,39 @@ +from typing import Protocol +from abc import abstractmethod + + +class MethodExample(Protocol): + def first(self) -> int: + return 42 + + @abstractmethod + def second(self) -> int: + raise NotImplementedError + + +class MethodExampleImpl1: + def first(self) -> int: + return 42 + + def second(self) -> int: + return 24 + + +class MethodExampleImpl2(MethodExample): + def second(self) -> int: + return 24 + + +class MethodExampleImpl3: + def second(self) -> int: + return 24 + + +def example(e: MethodExample) -> None: + print(e.first()) + print(e.second()) + + +example(MethodExampleImpl1()) +example(MethodExampleImpl2()) +example(MethodExampleImpl3()) diff --git a/python/testData/inspections/PyTypeCheckerInspection/AgainstTypingProtocolWithImplementedVariable.py b/python/testData/inspections/PyTypeCheckerInspection/AgainstTypingProtocolWithImplementedVariable.py new file mode 100644 index 000000000000..f66d07a822f7 --- /dev/null +++ b/python/testData/inspections/PyTypeCheckerInspection/AgainstTypingProtocolWithImplementedVariable.py @@ -0,0 +1,32 @@ +from typing import Protocol + + +class VariableExample(Protocol): + name: str + value: int = 0 + + +class VariableExampleImpl1: + def __init__(self, name: str, value: int) -> None: + self.name = name + self.value = value + + +class VariableExampleImpl2(VariableExample): + def __init__(self, name: str) -> None: + self.name = name + + +class VariableExampleImpl3: + def __init__(self, name: str) -> None: + self.name = name + + +def example(e: VariableExample) -> None: + print(e.name) + print(e.value) + + +example(VariableExampleImpl1("1", 1)) +example(VariableExampleImpl2("1")) +example(VariableExampleImpl3("1")) diff --git a/python/testData/inspections/PyTypeCheckerInspection/AgainstTypingProtocolWrongTypes.py b/python/testData/inspections/PyTypeCheckerInspection/AgainstTypingProtocolWrongTypes.py new file mode 100644 index 000000000000..e80b3d59ec32 --- /dev/null +++ b/python/testData/inspections/PyTypeCheckerInspection/AgainstTypingProtocolWrongTypes.py @@ -0,0 +1,43 @@ +from typing import Protocol + + +class MyProtocol(Protocol): + attr: int + def func(self, p: int) -> str: + pass + + +class MyClass1: + def __init__(self, attr: int) -> None: + self.attr = attr + + + def func(self, p: str) -> int: + pass + + +class MyClass2: + def __init__(self, attr: str) -> None: + self.attr = attr + + + def func(self, p: int) -> str: + pass + + +class MyClass3: + def __init__(self, attr: str) -> None: + self.attr = attr + + + def func(self, p: str) -> int: + pass + + +def foo(m: MyProtocol): + pass + + +foo(MyClass1(1)) +foo(MyClass2("1")) +foo(MyClass3("1")) \ No newline at end of file diff --git a/python/testData/inspections/PyTypeCheckerInspection/TypingProtocolAgainstProtocol.py b/python/testData/inspections/PyTypeCheckerInspection/TypingProtocolAgainstProtocol.py new file mode 100644 index 000000000000..0bf76079dbdd --- /dev/null +++ b/python/testData/inspections/PyTypeCheckerInspection/TypingProtocolAgainstProtocol.py @@ -0,0 +1,41 @@ +from typing import Protocol + + +class MyProtocol1(Protocol): + attr: int + + def func(self, p: int) -> str: + pass + + +class MyProtocol2(Protocol): + attr: int + more_attr: int + + def func(self, p: int) -> str: + pass + + def more_func(self, p: str) -> int: + pass + + +class MyProtocol3(Protocol): + attr: str + more_attr: str + + def func(self, p: str) -> int: + pass + + def more_func(self, p: int) -> str: + pass + + +def foo(p: MyProtocol1): + pass + + +v1: MyProtocol2 +v2: MyProtocol3 + +foo(v1) +foo(v2) diff --git a/python/testSrc/com/jetbrains/python/inspections/PyTypeCheckerInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/PyTypeCheckerInspectionTest.java index 7354ff89741c..7340798045da 100644 --- a/python/testSrc/com/jetbrains/python/inspections/PyTypeCheckerInspectionTest.java +++ b/python/testSrc/com/jetbrains/python/inspections/PyTypeCheckerInspectionTest.java @@ -470,4 +470,44 @@ public class PyTypeCheckerInspectionTest extends PyInspectionTestCase { public void testModuleWithGetAttr() { runWithLanguageLevel(LanguageLevel.PYTHON37, this::doMultiFileTest); } + + // PY-26628 + public void testAgainstTypingProtocol() { + runWithLanguageLevel(LanguageLevel.PYTHON35, this::doTest); + } + + // PY-26628 + public void testAgainstTypingProtocolWithImplementedMethod() { + runWithLanguageLevel(LanguageLevel.PYTHON35, this::doTest); + } + + // PY-26628 + public void testAgainstTypingProtocolWithImplementedVariable() { + runWithLanguageLevel(LanguageLevel.PYTHON36, this::doTest); + } + + // PY-26628 + public void testAgainstMergedTypingProtocols() { + runWithLanguageLevel(LanguageLevel.PYTHON35, this::doTest); + } + + // PY-26628 + public void testAgainstGenericTypingProtocol() { + runWithLanguageLevel(LanguageLevel.PYTHON36, this::doTest); + } + + // PY-26628 + public void testAgainstRecursiveTypingProtocol() { + runWithLanguageLevel(LanguageLevel.PYTHON35, this::doTest); + } + + // PY-26628 + public void testAgainstTypingProtocolWrongTypes() { + runWithLanguageLevel(LanguageLevel.PYTHON36, this::doTest); + } + + // PY-26628 + public void testTypingProtocolAgainstProtocol() { + runWithLanguageLevel(LanguageLevel.PYTHON36, this::doTest); + } }