From d7f199091081e072a3932216534c85b49fcc91cf Mon Sep 17 00:00:00 2001 From: Semyon Proshev Date: Thu, 13 Oct 2016 15:35:14 +0300 Subject: [PATCH] Update PyTypeChecker to pass PyTypeTest.testDictFromTuple. If actual is union of tuples and expected is tuple then convert actual to tuple of unions and match it --- .../python/psi/types/PyTypeChecker.java | 56 ++++++++++++++++++- 1 file changed, 55 insertions(+), 1 deletion(-) diff --git a/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java b/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java index 8d4f3255182f..d6862d48bbaa 100644 --- a/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java +++ b/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java @@ -25,6 +25,7 @@ import com.jetbrains.python.psi.impl.PyBuiltinCache; import com.jetbrains.python.psi.impl.PyCallExpressionHelper; import com.jetbrains.python.psi.resolve.PyResolveContext; import com.jetbrains.python.psi.resolve.RatedResolveResult; +import one.util.streamex.StreamEx; import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; @@ -99,7 +100,18 @@ public class PyTypeChecker { return true; } if (actual instanceof PyUnionType) { - for (PyType m : ((PyUnionType)actual).getMembers()) { + final PyUnionType actualUnionType = (PyUnionType)actual; + + if (expected instanceof PyTupleType) { + final PyTupleType expectedTupleType = (PyTupleType)expected; + final int elementCount = expectedTupleType.getElementCount(); + + if (!expectedTupleType.isHomogeneous() && consistsOfSameElementNumberTuples(actualUnionType, elementCount)) { + return substituteExpectedElementsWithUnions(expectedTupleType, elementCount, actualUnionType, context, substitutions, recursive); + } + } + + for (PyType m : actualUnionType.getMembers()) { if (match(expected, m, context, substitutions, recursive)) { return true; } @@ -244,6 +256,48 @@ public class PyTypeChecker { return matchNumericTypes(expected, actual); } + private static boolean consistsOfSameElementNumberTuples(@NotNull PyUnionType unionType, int elementCount) { + for (PyType type : unionType.getMembers()) { + if (type instanceof PyTupleType) { + final PyTupleType tupleType = (PyTupleType)type; + + if (!tupleType.isHomogeneous() && elementCount != tupleType.getElementCount()) { + return false; + } + } + else { + return false; + } + } + + return true; + } + + private static boolean substituteExpectedElementsWithUnions(@NotNull PyTupleType expected, + int elementCount, + @NotNull PyUnionType actual, + @NotNull TypeEvalContext context, + @Nullable Map substitutions, + boolean recursive) { + for (int i = 0; i < elementCount; i++) { + final int currentIndex = i; + + final PyType elementType = PyUnionType.union( + StreamEx + .of(actual.getMembers()) + .select(PyTupleType.class) + .map(type -> type.getElementType(currentIndex)) + .toList() + ); + + if (!match(expected.getElementType(i), elementType, context, substitutions, recursive)) { + return false; + } + } + + return true; + } + private static boolean matchNumericTypes(PyType expected, PyType actual) { final String superName = expected.getName(); final String subName = actual.getName();