From 1515dae7dfe37ee6051a0855e9dd861c352861df Mon Sep 17 00:00:00 2001 From: Andrey Vlasovskikh Date: Fri, 14 Apr 2017 22:40:26 +0300 Subject: [PATCH] Refactored analyzing call arguments for type checking --- .../inspections/PyTypeCheckerInspection.java | 55 ++++-------------- .../psi/impl/PyCallExpressionHelper.java | 37 ++++++++++++ .../python/psi/types/PyTypeChecker.java | 56 ++++++++----------- 3 files changed, 70 insertions(+), 78 deletions(-) diff --git a/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java b/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java index c6ae967180f1..dc85724e0ed9 100644 --- a/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java +++ b/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java @@ -213,16 +213,10 @@ public class PyTypeCheckerInspection extends PyInspection { @NotNull private AnalyzeCalleeResults analyzeCallee(@NotNull PyTypeChecker.AnalyzeCallResults results) { final List result = new ArrayList<>(); - final PyExpression receiver = results.getReceiver(); - - Map substitutions = null; + final Map substitutions = PyTypeChecker.unifyReceiver(results.getReceiver(), myTypeEvalContext); for (Map.Entry entry : results.getMapping().getMappedParameters().entrySet()) { - final AnalyzeArgumentResult argumentResult = - analyzeArgument(receiver, entry.getValue(), entry.getKey(), substitutions); - - substitutions = argumentResult.mySubstitutions; - + final AnalyzeArgumentResult argumentResult = analyzeArgument(entry.getValue(), entry.getKey(), substitutions); result.add(argumentResult); } @@ -247,41 +241,17 @@ public class PyTypeCheckerInspection extends PyInspection { .collect(Collectors.toList()); } - /** - * @param receiver call receiver - * @param parameter callee parameter - * @param argument passed argument - * @param substitutions generics substitutions - * @return an object that contains expected argument type, expected argument type after substitution, actual argument type, - * flag with result of matching actual type against expected one and generics substitutions - * Note: generics substitutions are not recalculated if they were calculated before - */ @NotNull - private AnalyzeArgumentResult analyzeArgument(@Nullable PyExpression receiver, - @NotNull PyNamedParameter parameter, + private AnalyzeArgumentResult analyzeArgument(@NotNull PyNamedParameter parameter, @NotNull PyExpression argument, - @Nullable Map substitutions) { + @NotNull Map substitutions) { final PyType expectedArgumentType = parameter.getArgumentType(myTypeEvalContext); final PyType actualArgumentType = myTypeEvalContext.getType(argument); - - if (expectedArgumentType == null) { - return new AnalyzeArgumentResult(argument, null, null, actualArgumentType, true, substitutions); - } - else { - substitutions = substitutions != null ? substitutions : PyTypeChecker.unifyReceiver(receiver, myTypeEvalContext); - - final PyType expectedTypeAfterSubstitution = PyTypeChecker.hasGenerics(expectedArgumentType, myTypeEvalContext) - ? PyTypeChecker.substitute(expectedArgumentType, substitutions, myTypeEvalContext) - : null; - return new AnalyzeArgumentResult( - argument, - expectedArgumentType, - expectedTypeAfterSubstitution, - actualArgumentType, - actualArgumentType == null || PyTypeChecker.match(expectedArgumentType, actualArgumentType, myTypeEvalContext, substitutions), - substitutions - ); - } + 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); } } @@ -344,21 +314,16 @@ public class PyTypeCheckerInspection extends PyInspection { private final boolean myIsMatched; - @Nullable - private final Map mySubstitutions; - public AnalyzeArgumentResult(@NotNull PyExpression argument, @Nullable PyType expectedType, @Nullable PyType expectedTypeAfterSubstitution, @Nullable PyType actualType, - boolean isMatched, - @Nullable Map substitutions) { + boolean isMatched) { myArgument = argument; myExpectedType = expectedType; myExpectedTypeAfterSubstitution = expectedTypeAfterSubstitution; myActualType = actualType; myIsMatched = isMatched; - mySubstitutions = substitutions; } @NotNull diff --git a/python/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.java b/python/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.java index 6ae39240dac9..6bd5056db166 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.java +++ b/python/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.java @@ -663,6 +663,43 @@ public class PyCallExpressionHelper { return mapArguments(callSite, callable, parameters, context); } + @NotNull + public static List getArgumentsMappedToPositionalContainer(@NotNull Map mapping) { + return mapping.entrySet().stream() + .filter(e -> e.getValue().isPositionalContainer()) + .map(e -> e.getKey()).collect(Collectors.toList()); + } + + @NotNull + public static List getArgumentsMappedToKeywordContainer(@NotNull Map mapping) { + return mapping.entrySet().stream() + .filter(e -> e.getValue().isKeywordContainer()) + .map(e -> e.getKey()).collect(Collectors.toList()); + } + + @NotNull + public static Map getRegularMappedParameters(@NotNull Map mapping) { + final Map result = new LinkedHashMap<>(); + for (Map.Entry entry : mapping.entrySet()) { + final PyExpression argument = entry.getKey(); + final PyNamedParameter parameter = entry.getValue(); + if (!parameter.isPositionalContainer() && !parameter.isKeywordContainer()) { + result.put(argument, parameter); + } + } + return result; + } + + @Nullable + public static PyNamedParameter getMappedPositionalContainer(@NotNull Map mapping) { + return mapping.values().stream().filter(p -> p.isPositionalContainer()).findFirst().orElse(null); + } + + @Nullable + public static PyNamedParameter getMappedKeywordContainer(@NotNull Map mapping) { + return mapping.values().stream().filter(p -> p.isKeywordContainer()).findFirst().orElse(null); + } + @NotNull private static ArgumentMappingResults analyzeArguments(@NotNull List arguments, @NotNull List parameters) { boolean seenSingleStar = false; diff --git a/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java b/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java index d9ec0d56cacf..8ed54c187f3b 100644 --- a/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java +++ b/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java @@ -23,7 +23,6 @@ import com.intellij.util.containers.ContainerUtil; import com.jetbrains.python.PyNames; import com.jetbrains.python.psi.*; import com.jetbrains.python.psi.impl.PyBuiltinCache; -import com.jetbrains.python.psi.impl.PyCallExpressionHelper; import com.jetbrains.python.psi.impl.PyTypeProvider; import com.jetbrains.python.psi.resolve.PyResolveContext; import com.jetbrains.python.psi.resolve.RatedResolveResult; @@ -32,8 +31,10 @@ import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; 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 @@ -519,43 +520,32 @@ 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()) { + for (Map.Entry entry : getRegularMappedParameters(arguments).entrySet()) { + final PyType argumentType = context.getType(entry.getKey()); final PyNamedParameter parameter = entry.getValue(); - final PyType actualArgType = context.getType(entry.getKey()); - - if (parameter.isPositionalContainer()) { - if (positionalParameter == null) positionalParameter = parameter; - positionalTypes.add(actualArgType); - } - else if (parameter.isKeywordContainer()) { - if (keywordParameter == null) keywordParameter = parameter; - keywordTypes.add(actualArgType); - } - else if (!match(parameter.getArgumentType(context), actualArgType, context, substitutions)) { + if (!match(parameter.getArgumentType(context), argumentType, context, substitutions)) { return null; } } - - if (positionalParameter != null && - !match(positionalParameter.getArgumentType(context), PyUnionType.union(positionalTypes), context, substitutions)) { + if (!matchContainer(getMappedPositionalContainer(arguments), getArgumentsMappedToPositionalContainer(arguments), substitutions, + context)) { return null; } - - if (keywordParameter != null && - !match(keywordParameter.getArgumentType(context), PyUnionType.union(keywordTypes), context, substitutions)) { + if (!matchContainer(getMappedKeywordContainer(arguments), getArgumentsMappedToKeywordContainer(arguments), substitutions, context)) { return null; } - return substitutions; } + private static boolean matchContainer(@Nullable PyNamedParameter container, @NotNull List arguments, + @NotNull Map substitutions, @NotNull TypeEvalContext context) { + if (container == null) { + return true; + } + final List types = arguments.stream().map(e -> context.getType(e)).collect(Collectors.toList()); + return match(container.getArgumentType(context), PyUnionType.union(types), context, substitutions); + } + @NotNull public static Map unifyReceiver(@Nullable PyExpression receiver, @NotNull TypeEvalContext context) { final Map substitutions = new LinkedHashMap<>(); @@ -624,7 +614,7 @@ public class PyTypeChecker { final List results = new ArrayList<>(); for (PyCallable callable : multiResolveCallee(callSite, context)) { final PyExpression receiver = getReceiver(callSite, callable); - final PyCallExpressionHelper.ArgumentMappingResults mapping = PyCallExpressionHelper.mapArguments(callSite, callable, context); + final ArgumentMappingResults mapping = mapArguments(callSite, callable, context); results.add(new AnalyzeCallResults(callable, receiver, mapping)); } return results; @@ -760,8 +750,8 @@ public class PyTypeChecker { final PyCallExpression callExpr = (PyCallExpression)callSite; final PyExpression callee = callExpr.getCallee(); if (callee instanceof PyReferenceExpression && callable instanceof PyFunction) { - implicitOffset = PyCallExpressionHelper.getImplicitArgumentCount((PyReferenceExpression)callee, (PyFunction)callable, - resolveContext); + implicitOffset = getImplicitArgumentCount((PyReferenceExpression)callee, (PyFunction)callable, + resolveContext); } else { implicitOffset = 0; @@ -822,10 +812,10 @@ public class PyTypeChecker { public static class AnalyzeCallResults { @NotNull private final PyCallable myCallable; @Nullable private final PyExpression myReceiver; - @NotNull private final PyCallExpressionHelper.ArgumentMappingResults myMapping; + @NotNull private final ArgumentMappingResults myMapping; public AnalyzeCallResults(@NotNull PyCallable callable, @Nullable PyExpression receiver, - @NotNull PyCallExpressionHelper.ArgumentMappingResults mapping) { + @NotNull ArgumentMappingResults mapping) { myCallable = callable; myReceiver = receiver; myMapping = mapping; @@ -842,7 +832,7 @@ public class PyTypeChecker { } @NotNull - public PyCallExpressionHelper.ArgumentMappingResults getMapping() { + public ArgumentMappingResults getMapping() { return myMapping; } }