Implement matching types against protocol definitions (PY-26628)

This commit is contained in:
Semyon Proshev
2018-02-05 20:58:29 +03:00
parent 5eecc4b814
commit 6fa83614d4
5 changed files with 91 additions and 3 deletions
@@ -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<PyExpression, PyCallableParameter> arguments =
@@ -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<PyGenericType, PyType> substitutions) {
return PyTypeChecker.match(parameterType, argumentType, myTypeEvalContext, substitutions) &&
!PyTypingTypeProvider.matchingProtocolDefinitions(parameterType, argumentType, myTypeEvalContext);
}
@Nullable
private PyType substituteGenerics(@Nullable PyType expectedArgumentType, @NotNull Map<PyGenericType, PyType> substitutions) {
return PyTypeChecker.hasGenerics(expectedArgumentType, myTypeEvalContext)
@@ -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);
}
@@ -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(<warning descr="Only concrete class can be given where 'Type[Proto]' is expected">Proto</warning>)
foo(Concrete1)
foo(Concrete2)
foo(<warning descr="Expected type 'Type[Proto]', got 'Type[Concrete3]' instead">Concrete3</warning>)
foo(Concrete4) # matched as inheritor
foo(<warning descr="Only concrete class can be given where 'Type[Proto]' is expected">NewProto</warning>)
bar(<warning descr="Only concrete class can be given where 'Type[Proto]' is expected">Proto</warning>)
bar(Concrete1)
bar(Concrete2)
bar(<warning descr="Expected type 'Type[Proto]', got 'Type[Concrete3]' instead">Concrete3</warning>)
bar(Concrete4) # matched as inheritor
bar(<warning descr="Only concrete class can be given where 'Type[Proto]' is expected">NewProto</warning>)
@@ -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);
}
}