Implement matching types against protocols (PY-26628)

Update PyTypeChecker to prevent recursion while matching types.
SOE was received while matching `self` parameters for methods inside protocol and other class.
This commit is contained in:
Semyon Proshev
2018-02-05 20:58:29 +03:00
parent 77e5e45220
commit 5eecc4b814
11 changed files with 414 additions and 33 deletions
@@ -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<PyExpression, PyCallableParameter> arguments =
@@ -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<PyGenericType, PyType> 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<PyGenericType, PyType> substitutions, boolean recursive) {
@Nullable Map<PyGenericType, PyType> substitutions, boolean recursive,
@NotNull Set<Pair<PyType, PyType>> matching) {
final Pair<PyType, PyType> 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<PyGenericType, PyType> substitutions,
boolean recursive,
@NotNull Set<Pair<PyType, PyType>> 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<PyGenericType> 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<? 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));
}
}
}
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<PyType> superElementTypes = ((PyCollectionType)expected).getElementTypes();
final PyCollectionType actualCollectionType = as(actual, PyCollectionType.class);
final List<PyType> 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<PyGenericType, PyType> substitutions,
boolean recursive) {
boolean recursive,
@NotNull Set<Pair<PyType, PyType>> 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<PyGenericType, PyType> substitutions,
boolean recursive,
@NotNull Set<Pair<PyType, PyType>> matching) {
// TODO: Match generic parameters based on the correspondence between the generic parameters of subClass and its base classes
final List<PyType> superElementTypes = expected.getElementTypes();
final PyCollectionType actualCollectionType = as(actual, PyCollectionType.class);
final List<PyType> 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();
@@ -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)
@@ -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(<warning descr="Expected type 'SizedAndClosable', got 'int' instead">1</warning>)
@@ -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())
@@ -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(<warning descr="Expected type 'SupportsClose', got 'Type[Resource]' instead">Resource</warning>)
close(<warning descr="Expected type 'SupportsClose', got 'int' instead">1</warning>)
@@ -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(<warning descr="Expected type 'MethodExample', got 'MethodExampleImpl3' instead">MethodExampleImpl3()</warning>)
@@ -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(<warning descr="Expected type 'VariableExample', got 'VariableExampleImpl3' instead">VariableExampleImpl3("1")</warning>)
@@ -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(<warning descr="Expected type 'MyProtocol', got 'MyClass1' instead">MyClass1(1)</warning>)
foo(<warning descr="Expected type 'MyProtocol', got 'MyClass2' instead">MyClass2("1")</warning>)
foo(<warning descr="Expected type 'MyProtocol', got 'MyClass3' instead">MyClass3("1")</warning>)
@@ -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(<warning descr="Expected type 'MyProtocol1', got 'MyProtocol3' instead">v2</warning>)
@@ -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);
}
}