From 45ec95d9392ddd97b0fd951d038a169a99fcf7d8 Mon Sep 17 00:00:00 2001 From: Semyon Proshev Date: Tue, 22 Nov 2016 21:32:33 +0300 Subject: [PATCH] PY-19723 Fixed: Type hinting of arbitrary argument lists and default argument values Update PyTypeChecker to substitute type vars with types of positional and keyword args (incl. heterogeneous ones) --- .../inspections/PyTypeCheckerInspection.java | 21 +------- .../python/psi/types/PyTypeChecker.java | 53 ++++++++++++++++--- .../com/jetbrains/python/PyTypeTest.java | 48 +++++++++++++++++ 3 files changed, 96 insertions(+), 26 deletions(-) diff --git a/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java b/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java index 267c0da85310..a49eb8027902 100644 --- a/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java +++ b/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java @@ -24,7 +24,6 @@ import com.intellij.openapi.util.Key; import com.intellij.openapi.util.Pair; import com.intellij.openapi.util.text.StringUtil; import com.intellij.psi.PsiElementVisitor; -import com.intellij.util.containers.ContainerUtil; import com.intellij.util.containers.hash.LinkedHashMap; import com.jetbrains.python.PyNames; import com.jetbrains.python.codeInsight.controlflow.ScopeOwner; @@ -204,7 +203,7 @@ public class PyTypeCheckerInspection extends PyInspection { for (Map.Entry entry : mapping.entrySet()) { final PyNamedParameter param = entry.getValue(); final PyExpression arg = entry.getKey(); - final PyType expectedArgType = getExpectedArgumentType(param); + final PyType expectedArgType = PyTypeChecker.getExpectedArgumentType(param, myTypeEvalContext); if (expectedArgType == null) { continue; } @@ -221,24 +220,6 @@ public class PyTypeCheckerInspection extends PyInspection { return problems; } - @Nullable - private PyType getExpectedArgumentType(@NotNull PyNamedParameter parameter) { - final PyType parameterType = myTypeEvalContext.getType(parameter); - - if (parameterType instanceof PyCollectionType) { - final PyCollectionType paramCollectionType = (PyCollectionType)parameterType; - - if (parameter.isPositionalContainer()) { - return paramCollectionType.getIteratedItemType(); - } - else if (parameter.isKeywordContainer()) { - return ContainerUtil.getOrElse(paramCollectionType.getElementTypes(myTypeEvalContext), 1, null); - } - } - - return parameterType; - } - @Nullable private static Pair checkTypes(@Nullable PyType expected, @Nullable PyType actual, diff --git a/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java b/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java index 0a315c4813a3..f214cfe5e6f5 100644 --- a/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java +++ b/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java @@ -468,17 +468,40 @@ public class PyTypeChecker { @NotNull Map arguments, @NotNull TypeEvalContext context) { final Map substitutions = unifyReceiver(receiver, context); + + PyNamedParameter positionalParameter = null; + final List positionalTypes = new ArrayList<>(); + + PyNamedParameter keywordParameter = null; + final List keywordTypes = new ArrayList<>(); + for (Map.Entry entry : arguments.entrySet()) { - final PyNamedParameter p = entry.getValue(); - if (p.isPositionalContainer() || p.isKeywordContainer()) { - continue; + final PyNamedParameter parameter = entry.getValue(); + final PyType actualArgType = context.getType(entry.getKey()); + + if (parameter.isPositionalContainer()) { + if (positionalParameter == null) positionalParameter = parameter; + positionalTypes.add(actualArgType); } - final PyType argType = context.getType(entry.getKey()); - final PyType paramType = context.getType(p); - if (!match(paramType, argType, context, substitutions)) { + else if (parameter.isKeywordContainer()) { + if (keywordParameter == null) keywordParameter = parameter; + keywordTypes.add(actualArgType); + } + else if (!match(getExpectedArgumentType(parameter, context), actualArgType, context, substitutions)) { return null; } } + + if (positionalParameter != null && + !match(getExpectedArgumentType(positionalParameter, context), PyUnionType.union(positionalTypes), context, substitutions)) { + return null; + } + + if (keywordParameter != null && + !match(getExpectedArgumentType(keywordParameter, context), PyUnionType.union(keywordTypes), context, substitutions)) { + return null; + } + return substitutions; } @@ -734,6 +757,24 @@ public class PyTypeChecker { } } + @Nullable + public static PyType getExpectedArgumentType(@NotNull PyNamedParameter parameter, @NotNull TypeEvalContext context) { + final PyType parameterType = context.getType(parameter); + + if (parameterType instanceof PyCollectionType) { + final PyCollectionType paramCollectionType = (PyCollectionType)parameterType; + + if (parameter.isPositionalContainer()) { + return paramCollectionType.getIteratedItemType(); + } + else if (parameter.isKeywordContainer()) { + return ContainerUtil.getOrElse(paramCollectionType.getElementTypes(context), 1, null); + } + } + + return parameterType; + } + public static class AnalyzeCallResults { @NotNull private final PyCallable myCallable; @Nullable private final PyExpression myReceiver; diff --git a/python/testSrc/com/jetbrains/python/PyTypeTest.java b/python/testSrc/com/jetbrains/python/PyTypeTest.java index dfa1371b9f0f..7834ff6adb81 100644 --- a/python/testSrc/com/jetbrains/python/PyTypeTest.java +++ b/python/testSrc/com/jetbrains/python/PyTypeTest.java @@ -1456,6 +1456,54 @@ public class PyTypeTest extends PyTestCase { " pass"); } + // PY-19723 + public void testTypeVarSubstitutionInPositionalArgs() { + doTest("int", + "def foo(*args):" + + " \"\"\"\n" + + " :type args: T\n" + + " :rtype: T\n" + + " \"\"\"\n" + + " pass\n" + + "expr = foo(1)"); + } + + // PY-19723 + public void testTypeVarSubstitutionInHeterogeneousPositionalArgs() { + doTest("Union[int, str]", + "def foo(*args):" + + " \"\"\"\n" + + " :type args: T\n" + + " :rtype: T\n" + + " \"\"\"\n" + + " pass\n" + + "expr = foo(1, \"2\")"); + } + + // PY-19723 + public void testTypeVarSubstitutionInKeywordArgs() { + doTest("int", + "def foo(**kwargs):" + + " \"\"\"\n" + + " :type kwargs: T\n" + + " :rtype: T\n" + + " \"\"\"\n" + + " pass\n" + + "expr = foo(a=1)"); + } + + // PY-19723 + public void testTypeVarSubstitutionInHeterogeneousKeywordArgs() { + doTest("Union[int, str]", + "def foo(**kwargs):" + + " \"\"\"\n" + + " :type kwargs: T\n" + + " :rtype: T\n" + + " \"\"\"\n" + + " pass\n" + + "expr = foo(a=1, b=\"2\")"); + } + private static List getTypeEvalContexts(@NotNull PyExpression element) { return ImmutableList.of(TypeEvalContext.codeAnalysis(element.getProject(), element.getContainingFile()).withTracing(), TypeEvalContext.userInitiated(element.getProject(), element.getContainingFile()).withTracing());