diff --git a/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java b/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java index dc85724e0ed9..2b06ebf7eaa0 100644 --- a/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java +++ b/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java @@ -28,7 +28,6 @@ import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider; import com.jetbrains.python.documentation.PythonDocumentationProvider; import com.jetbrains.python.inspections.quickfix.PyMakeFunctionReturnTypeQuickFix; import com.jetbrains.python.psi.*; -import com.jetbrains.python.psi.impl.PyCallExpressionHelper; import com.jetbrains.python.psi.types.*; import one.util.streamex.StreamEx; import org.jetbrains.annotations.Nls; @@ -39,6 +38,7 @@ import java.util.*; import java.util.stream.Collectors; import static com.jetbrains.python.psi.PyUtil.as; +import static com.jetbrains.python.psi.impl.PyCallExpressionHelper.*; /** * @author vlan @@ -206,7 +206,7 @@ public class PyTypeCheckerInspection extends PyInspection { } private static boolean callDoesNotHaveUnmappedArgumentsAndUnfilledParameters(@NotNull PyTypeChecker.AnalyzeCallResults callResults) { - final PyCallExpressionHelper.ArgumentMappingResults mapping = callResults.getMapping(); + final ArgumentMappingResults mapping = callResults.getMapping(); return mapping.getUnmappedArguments().isEmpty() && mapping.getUnmappedParameters().isEmpty(); } @@ -214,15 +214,57 @@ public class PyTypeCheckerInspection extends PyInspection { private AnalyzeCalleeResults analyzeCallee(@NotNull PyTypeChecker.AnalyzeCallResults results) { final List result = new ArrayList<>(); final Map substitutions = PyTypeChecker.unifyReceiver(results.getReceiver(), myTypeEvalContext); - - for (Map.Entry entry : results.getMapping().getMappedParameters().entrySet()) { - final AnalyzeArgumentResult argumentResult = analyzeArgument(entry.getValue(), entry.getKey(), substitutions); - result.add(argumentResult); + final Map mapping = results.getMapping().getMappedParameters(); + for (Map.Entry entry : getRegularMappedParameters(mapping).entrySet()) { + final PyExpression argument = entry.getKey(); + final PyNamedParameter parameter = entry.getValue(); + final PyType expected = parameter.getArgumentType(myTypeEvalContext); + final PyType actual = myTypeEvalContext.getType(argument); + final boolean matched = PyTypeChecker.match(expected, actual, myTypeEvalContext, substitutions); + result.add(new AnalyzeArgumentResult(argument, expected, substituteGenerics(expected, substitutions), actual, matched)); + } + final PyNamedParameter positionalContainer = getMappedPositionalContainer(mapping); + if (positionalContainer != null) { + result.addAll(analyzeContainerMapping(positionalContainer, getArgumentsMappedToPositionalContainer(mapping), substitutions)); + } + final PyNamedParameter keywordContainer = getMappedKeywordContainer(mapping); + if (keywordContainer != null) { + result.addAll(analyzeContainerMapping(keywordContainer, getArgumentsMappedToKeywordContainer(mapping), substitutions)); } - return new AnalyzeCalleeResults(results.getCallable(), result); } + @NotNull + private List analyzeContainerMapping(@NotNull PyNamedParameter container, @NotNull List arguments, + @NotNull Map substitutions) { + final PyType expected = container.getArgumentType(myTypeEvalContext); + final PyType expectedWithSubstitutions = substituteGenerics(expected, substitutions); + // 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); + return arguments.stream() + .map(argument -> new AnalyzeArgumentResult(argument, expected, expectedWithSubstitutions, actual, matched)) + .collect(Collectors.toList()); + } + else { + return arguments.stream() + .map(argument -> { + final PyType actual = myTypeEvalContext.getType(argument); + final boolean matched = PyTypeChecker.match(expected, actual, myTypeEvalContext, substitutions); + return new AnalyzeArgumentResult(argument, expected, expectedWithSubstitutions, actual, matched); + }) + .collect(Collectors.toList()); + } + } + + @Nullable + private PyType substituteGenerics(@Nullable PyType expectedArgumentType, @NotNull Map substitutions) { + return PyTypeChecker.hasGenerics(expectedArgumentType, myTypeEvalContext) + ? PyTypeChecker.substitute(expectedArgumentType, substitutions, myTypeEvalContext) + : null; + } + private static boolean matchedCalleeResultsExist(@NotNull List calleesResults) { return calleesResults .stream() @@ -240,19 +282,6 @@ public class PyTypeCheckerInspection extends PyInspection { .map(AnalyzeArgumentResult::getActualType) .collect(Collectors.toList()); } - - @NotNull - private AnalyzeArgumentResult analyzeArgument(@NotNull PyNamedParameter parameter, - @NotNull PyExpression argument, - @NotNull Map substitutions) { - final PyType expectedArgumentType = parameter.getArgumentType(myTypeEvalContext); - final PyType actualArgumentType = myTypeEvalContext.getType(argument); - final PyType expectedTypeAfterSubstitution = PyTypeChecker.hasGenerics(expectedArgumentType, myTypeEvalContext) - ? PyTypeChecker.substitute(expectedArgumentType, substitutions, myTypeEvalContext) - : null; - final boolean isMatched = PyTypeChecker.match(expectedArgumentType, actualArgumentType, myTypeEvalContext, substitutions); - return new AnalyzeArgumentResult(argument, expectedArgumentType, expectedTypeAfterSubstitution, actualArgumentType, isMatched); - } } @Override diff --git a/python/testData/inspections/PyTypeCheckerInspection/GenericKwargs.py b/python/testData/inspections/PyTypeCheckerInspection/GenericKwargs.py new file mode 100644 index 000000000000..019a49df8282 --- /dev/null +++ b/python/testData/inspections/PyTypeCheckerInspection/GenericKwargs.py @@ -0,0 +1,11 @@ +from typing import Any, TypeVar + + +T = TypeVar('T') + + +def generic_kwargs(**kwargs: T) -> None: + pass + + +generic_kwargs(a=1, b='foo') diff --git a/python/testSrc/com/jetbrains/python/Py3TypeTest.java b/python/testSrc/com/jetbrains/python/Py3TypeTest.java index ff5dd8a79fe0..b7054eac2008 100644 --- a/python/testSrc/com/jetbrains/python/Py3TypeTest.java +++ b/python/testSrc/com/jetbrains/python/Py3TypeTest.java @@ -519,6 +519,19 @@ public class Py3TypeTest extends PyTestCase { " print(expr)"); } + // PY-22513 + public void testGenericKwargs() { + doTest("Dict[str, Union[int, str]]", + "from typing import Any, Dict, TypeVar\n" + + "\n" + + "T = TypeVar('T')\n" + + "\n" + + "def generic_kwargs(**kwargs: T) -> Dict[str, T]:\n" + + " pass\n" + + "\n" + + "expr = generic_kwargs(a=1, b='foo')\n"); + } + private void doTest(final String expectedType, final String text) { myFixture.configureByText(PythonFileType.INSTANCE, text); final PyExpression expr = myFixture.findElementByText("expr", PyExpression.class); diff --git a/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java index a86c8023dfde..cbaef8617e4e 100644 --- a/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java +++ b/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java @@ -228,4 +228,9 @@ public class Py3TypeCheckerInspectionTest extends PyTestCase { public void testUnboundTypeVarsMatchClassObjectTypes() { doTest(); } + + // PY-22513 + public void testGenericKwargs() { + doTest(); + } }