Refactored analyzing call arguments for type checking

This commit is contained in:
Andrey Vlasovskikh
2017-04-15 00:17:04 +03:00
parent efead19dc6
commit 1515dae7df
3 changed files with 70 additions and 78 deletions
@@ -213,16 +213,10 @@ public class PyTypeCheckerInspection extends PyInspection {
@NotNull
private AnalyzeCalleeResults analyzeCallee(@NotNull PyTypeChecker.AnalyzeCallResults results) {
final List<AnalyzeArgumentResult> result = new ArrayList<>();
final PyExpression receiver = results.getReceiver();
Map<PyGenericType, PyType> substitutions = null;
final Map<PyGenericType, PyType> substitutions = PyTypeChecker.unifyReceiver(results.getReceiver(), myTypeEvalContext);
for (Map.Entry<PyExpression, PyNamedParameter> 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
* <i>Note: generics substitutions are not recalculated if they were calculated before</i>
*/
@NotNull
private AnalyzeArgumentResult analyzeArgument(@Nullable PyExpression receiver,
@NotNull PyNamedParameter parameter,
private AnalyzeArgumentResult analyzeArgument(@NotNull PyNamedParameter parameter,
@NotNull PyExpression argument,
@Nullable Map<PyGenericType, PyType> substitutions) {
@NotNull Map<PyGenericType, PyType> 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<PyGenericType, PyType> mySubstitutions;
public AnalyzeArgumentResult(@NotNull PyExpression argument,
@Nullable PyType expectedType,
@Nullable PyType expectedTypeAfterSubstitution,
@Nullable PyType actualType,
boolean isMatched,
@Nullable Map<PyGenericType, PyType> substitutions) {
boolean isMatched) {
myArgument = argument;
myExpectedType = expectedType;
myExpectedTypeAfterSubstitution = expectedTypeAfterSubstitution;
myActualType = actualType;
myIsMatched = isMatched;
mySubstitutions = substitutions;
}
@NotNull
@@ -663,6 +663,43 @@ public class PyCallExpressionHelper {
return mapArguments(callSite, callable, parameters, context);
}
@NotNull
public static List<PyExpression> getArgumentsMappedToPositionalContainer(@NotNull Map<PyExpression, PyNamedParameter> mapping) {
return mapping.entrySet().stream()
.filter(e -> e.getValue().isPositionalContainer())
.map(e -> e.getKey()).collect(Collectors.toList());
}
@NotNull
public static List<PyExpression> getArgumentsMappedToKeywordContainer(@NotNull Map<PyExpression, PyNamedParameter> mapping) {
return mapping.entrySet().stream()
.filter(e -> e.getValue().isKeywordContainer())
.map(e -> e.getKey()).collect(Collectors.toList());
}
@NotNull
public static Map<PyExpression, PyNamedParameter> getRegularMappedParameters(@NotNull Map<PyExpression, PyNamedParameter> mapping) {
final Map<PyExpression, PyNamedParameter> result = new LinkedHashMap<>();
for (Map.Entry<PyExpression, PyNamedParameter> 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<PyExpression, PyNamedParameter> mapping) {
return mapping.values().stream().filter(p -> p.isPositionalContainer()).findFirst().orElse(null);
}
@Nullable
public static PyNamedParameter getMappedKeywordContainer(@NotNull Map<PyExpression, PyNamedParameter> mapping) {
return mapping.values().stream().filter(p -> p.isKeywordContainer()).findFirst().orElse(null);
}
@NotNull
private static ArgumentMappingResults analyzeArguments(@NotNull List<PyExpression> arguments, @NotNull List<PyParameter> parameters) {
boolean seenSingleStar = false;
@@ -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<PyExpression, PyNamedParameter> arguments,
@NotNull TypeEvalContext context) {
final Map<PyGenericType, PyType> substitutions = unifyReceiver(receiver, context);
PyNamedParameter positionalParameter = null;
final List<PyType> positionalTypes = new ArrayList<>();
PyNamedParameter keywordParameter = null;
final List<PyType> keywordTypes = new ArrayList<>();
for (Map.Entry<PyExpression, PyNamedParameter> entry : arguments.entrySet()) {
for (Map.Entry<PyExpression, PyNamedParameter> 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<PyExpression> arguments,
@NotNull Map<PyGenericType, PyType> substitutions, @NotNull TypeEvalContext context) {
if (container == null) {
return true;
}
final List<PyType> 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<PyGenericType, PyType> unifyReceiver(@Nullable PyExpression receiver, @NotNull TypeEvalContext context) {
final Map<PyGenericType, PyType> substitutions = new LinkedHashMap<>();
@@ -624,7 +614,7 @@ public class PyTypeChecker {
final List<AnalyzeCallResults> 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;
}
}