Simpler implementation of mapping arguments to parameters

It's a part of rewriting CallArgumentsMapping in order to support
overloading functions in Python stubs. The old implementation was tied
to a single set of parameters whereas the new implementation could be
transformed to handle multiple sets of parameters, i.e. overloaded
signatures.
This commit is contained in:
Andrey Vlasovskikh
2015-08-19 14:47:11 +03:00
parent b4f977d75b
commit 8afbcce2b3
9 changed files with 213 additions and 25 deletions
@@ -23,6 +23,7 @@ import com.intellij.psi.PsiReference;
import com.jetbrains.python.PyBundle;
import com.jetbrains.python.inspections.quickfix.RemoveArgumentEqualDefaultQuickFix;
import com.jetbrains.python.psi.*;
import com.jetbrains.python.psi.impl.CallArgumentsMappingImpl;
import com.jetbrains.python.psi.impl.PyBuiltinCache;
import com.jetbrains.python.psi.types.PyClassType;
import org.jetbrains.annotations.Nls;
@@ -77,8 +78,7 @@ public class PyArgumentEqualDefaultInspection extends PyInspection {
if (func != null && hasSpecialCasedDefaults(func, node)) {
return;
}
CallArgumentsMapping result = list.analyzeCall(getResolveContext());
checkArguments(result, node.getArguments());
checkArguments(node, node.getArguments());
}
private static boolean hasSpecialCasedDefaults(PyCallable callable, PsiElement anchor) {
@@ -97,8 +97,8 @@ public class PyArgumentEqualDefaultInspection extends PyInspection {
return false;
}
private void checkArguments(CallArgumentsMapping result, PyExpression[] arguments) {
Map<PyExpression, PyNamedParameter> mapping = result.getPlainMappedParams();
private void checkArguments(PyCallExpression callExpr, PyExpression[] arguments) {
final Map<PyExpression, PyNamedParameter> mapping = CallArgumentsMappingImpl.map(callExpr, getResolveContext());
Set<PyExpression> problemElements = new HashSet<PyExpression>();
for (Map.Entry<PyExpression, PyNamedParameter> e : mapping.entrySet()) {
PyExpression defaultValue = e.getValue().getDefaultValue();
@@ -22,6 +22,7 @@ import com.intellij.psi.PsiElementVisitor;
import com.intellij.psi.util.PsiTreeUtil;
import com.jetbrains.python.PyBundle;
import com.jetbrains.python.psi.*;
import com.jetbrains.python.psi.impl.CallArgumentsMappingImpl;
import com.jetbrains.python.psi.types.PyClassType;
import com.jetbrains.python.psi.types.PyType;
import org.jetbrains.annotations.Nls;
@@ -90,13 +91,13 @@ public class PyCallByClassInspection extends PyInspection {
PyClass qual_class = qual_class_type.getPyClass();
final PyArgumentList arglist = call.getArgumentList();
if (arglist != null) {
CallArgumentsMapping analysis = arglist.analyzeCall(getResolveContext());
final PyCallExpression.PyMarkedCallee markedCallee = analysis.getMarkedCallee();
final PyCallExpression.PyMarkedCallee markedCallee = call.resolveCallee(getResolveContext());
if (markedCallee != null && markedCallee.getModifier() != STATICMETHOD) {
final List<PyParameter> params = PyUtil.getParameters(markedCallee.getCallable(), myTypeEvalContext);
if (params.size() > 0 && params.get(0) instanceof PyNamedParameter) {
PyNamedParameter first_param = (PyNamedParameter)params.get(0);
for (Map.Entry<PyExpression, PyNamedParameter> entry : analysis.getPlainMappedParams().entrySet()) {
final Map<PyExpression, PyNamedParameter> mapping = CallArgumentsMappingImpl.map(call, getResolveContext());
for (Map.Entry<PyExpression, PyNamedParameter> entry : mapping.entrySet()) {
// we ignore *arg and **arg which we cannot analyze
if (entry.getValue() == first_param) {
PyExpression first_arg = entry.getKey();
@@ -34,6 +34,7 @@ import com.jetbrains.python.PyNames;
import com.jetbrains.python.inspections.quickfix.PyUpdatePropertySignatureQuickFix;
import com.jetbrains.python.inspections.quickfix.RenameParameterQuickFix;
import com.jetbrains.python.psi.*;
import com.jetbrains.python.psi.impl.CallArgumentsMappingImpl;
import com.jetbrains.python.psi.impl.PyBuiltinCache;
import com.jetbrains.python.psi.types.PyClassType;
import com.jetbrains.python.psi.types.PyNoneType;
@@ -119,9 +120,9 @@ public class PyPropertyDefinitionInspection extends PyInspection {
assert call != null : "Property has a null call assigned to it";
final PyArgumentList arglist = call.getArgumentList();
assert arglist != null : "Property call has null arglist";
CallArgumentsMapping analysis = arglist.analyzeCall(getResolveContext());
// we assume fget, fset, fdel, doc names
for (Map.Entry<PyExpression, PyNamedParameter> entry : analysis.getPlainMappedParams().entrySet()) {
final Map<PyExpression, PyNamedParameter> mapping = CallArgumentsMappingImpl.map(call, getResolveContext());
for (Map.Entry<PyExpression, PyNamedParameter> entry : mapping.entrySet()) {
final String paramName = entry.getValue().getName();
PyExpression argument = PyUtil.peelArgument(entry.getKey());
checkPropertyCallArgument(paramName, argument, node.getContainingFile());
@@ -27,6 +27,7 @@ import com.jetbrains.python.PyBundle;
import com.jetbrains.python.documentation.PyDocumentationSettings;
import com.jetbrains.python.editor.PythonDocCommentUtil;
import com.jetbrains.python.psi.*;
import com.jetbrains.python.psi.impl.CallArgumentsMappingImpl;
import com.jetbrains.python.psi.resolve.PyResolveContext;
import com.jetbrains.python.refactoring.PyRefactoringUtil;
import org.jetbrains.annotations.NonNls;
@@ -63,8 +64,9 @@ public class PyRemoveParameterQuickFix implements LocalQuickFix {
if (callExpression instanceof PyCallExpression) {
final PyArgumentList argumentList = ((PyCallExpression)callExpression).getArgumentList();
if (argumentList != null) {
final CallArgumentsMapping mapping = argumentList.analyzeCall(PyResolveContext.noImplicits());
for (Map.Entry<PyExpression, PyNamedParameter> parameterEntry : mapping.getPlainMappedParams().entrySet()) {
final Map<PyExpression, PyNamedParameter> mapping = CallArgumentsMappingImpl.map((PyCallExpression)callExpression,
PyResolveContext.noImplicits());
for (Map.Entry<PyExpression, PyNamedParameter> parameterEntry : mapping.entrySet()) {
if (parameterEntry.getValue().equals(element)) {
parameterEntry.getKey().delete();
}
@@ -18,7 +18,10 @@ package com.jetbrains.python.psi.impl;
import com.intellij.psi.util.PsiTreeUtil;
import com.intellij.util.containers.ContainerUtil;
import com.jetbrains.python.psi.*;
import com.jetbrains.python.psi.types.*;
import com.jetbrains.python.psi.resolve.PyResolveContext;
import com.jetbrains.python.psi.types.PyTupleType;
import com.jetbrains.python.psi.types.PyType;
import com.jetbrains.python.psi.types.TypeEvalContext;
import org.jetbrains.annotations.NotNull;
import org.jetbrains.annotations.Nullable;
@@ -52,6 +55,180 @@ public class CallArgumentsMappingImpl implements CallArgumentsMapping {
myArgumentList = arglist;
}
@NotNull
public static Map<PyExpression, PyNamedParameter> map(@NotNull PyCallExpression callExpression,
@NotNull PyResolveContext resolveContext) {
return map(callExpression, 0, resolveContext);
}
@NotNull
public static Map<PyExpression, PyNamedParameter> map(@NotNull PyCallExpression callExpression,
int implicitArgumentOffset,
@NotNull PyResolveContext resolveContext) {
final Map<PyExpression, PyNamedParameter> results = new LinkedHashMap<PyExpression, PyNamedParameter>();
final PyArgumentList argumentList = callExpression.getArgumentList();
final PyCallExpression.PyMarkedCallee markedCallee = callExpression.resolveCallee(resolveContext, implicitArgumentOffset);
if (markedCallee != null && argumentList != null) {
final TypeEvalContext context = resolveContext.getTypeEvalContext();
final List<PyParameter> allParameters = PyUtil.getParameters(markedCallee.getCallable(), context);
final List<PyParameter> parameters = dropImplicitParameters(allParameters, markedCallee.getImplicitOffset());
final List<PyExpression> arguments = Arrays.asList(argumentList.getArguments());
final List<PyExpression> positionalArguments = filterPositionalArguments(arguments);
final List<PyKeywordArgument> keywordArguments = filterKeywordArguments(arguments);
final List<PyExpression> variadicPositionalArguments = filterVariadicPositionalArguments(arguments);
final List<PyExpression> variadicKeywordArguments = filterVariadicKeywordArguments(arguments);
boolean seenSingleStar = false;
final List<PyParameter> unmappedParameters = new ArrayList<PyParameter>();
for (PyParameter parameter : parameters) {
if (parameter instanceof PyNamedParameter) {
final PyNamedParameter namedParameter = (PyNamedParameter)parameter;
final String parameterName = namedParameter.getName();
if (namedParameter.isPositionalContainer()) {
if (variadicPositionalArguments.size() == 1) {
results.put(variadicPositionalArguments.remove(0), namedParameter);
}
else {
positionalArguments.clear();
variadicPositionalArguments.clear();
}
}
else if (namedParameter.isKeywordContainer()) {
if (variadicKeywordArguments.size() == 1) {
results.put(variadicKeywordArguments.remove(0), namedParameter);
}
else {
keywordArguments.clear();
variadicKeywordArguments.clear();
}
}
else if (seenSingleStar) {
final PyExpression keywordArgument = removeKeywordArgument(keywordArguments, parameterName);
if (keywordArgument != null) {
results.put(keywordArgument, namedParameter);
}
else if (variadicKeywordArguments.isEmpty()) {
unmappedParameters.add(namedParameter);
}
}
else {
if (!positionalArguments.isEmpty()) {
final PyExpression positionalArgument = next(positionalArguments);
if (positionalArgument != null) {
results.put(positionalArgument, namedParameter);
}
else {
unmappedParameters.add(namedParameter);
}
}
else {
final PyKeywordArgument keywordArgument = removeKeywordArgument(keywordArguments, parameterName);
if (keywordArgument != null) {
results.put(keywordArgument, namedParameter);
}
else if (variadicPositionalArguments.isEmpty() || variadicKeywordArguments.isEmpty()) {
unmappedParameters.add(namedParameter);
}
}
}
}
else if (parameter instanceof PyTupleParameter) {
unmappedParameters.add(parameter);
}
else if (parameter instanceof PySingleStarParameter) {
seenSingleStar = true;
}
else {
unmappedParameters.add(parameter);
}
}
}
return results;
}
@Nullable
private static PyKeywordArgument removeKeywordArgument(@NotNull List<PyKeywordArgument> arguments, @Nullable String name) {
PyKeywordArgument result = null;
for (PyKeywordArgument argument : arguments) {
final String keyword = argument.getKeyword();
if (keyword != null && keyword.equals(name)) {
result = argument;
break;
}
}
if (result != null) {
arguments.remove(result);
}
return result;
}
@NotNull
private static List<PyExpression> filterPositionalArguments(@NotNull List<PyExpression> arguments) {
final List<PyExpression> results = new ArrayList<PyExpression>();
for (PyExpression argument : arguments) {
if (isPositionalArg(argument)) {
results.add(argument);
}
}
return results;
}
@NotNull
private static List<PyKeywordArgument> filterKeywordArguments(@NotNull List<PyExpression> arguments) {
final List<PyKeywordArgument> results = new ArrayList<PyKeywordArgument>();
for (PyExpression argument : arguments) {
if (argument instanceof PyKeywordArgument) {
results.add((PyKeywordArgument)argument);
}
}
return results;
}
@NotNull
private static List<PyExpression> filterVariadicPositionalArguments(@NotNull List<PyExpression> arguments) {
final List<PyExpression> results = new ArrayList<PyExpression>();
for (PyExpression argument : arguments) {
if (argument != null && isVariadicPositionalArgument(argument)) {
results.add(argument);
}
}
return results;
}
@NotNull
private static List<PyExpression> filterVariadicKeywordArguments(@NotNull List<PyExpression> arguments) {
final List<PyExpression> results = new ArrayList<PyExpression>();
for (PyExpression argument : arguments) {
if (argument != null && isVariadicKeywordArgument(argument)) {
results.add(argument);
}
}
return results;
}
private static boolean isVariadicKeywordArgument(@NotNull PyExpression argument) {
return argument instanceof PyStarArgument && ((PyStarArgument)argument).isKeyword();
}
private static boolean isVariadicPositionalArgument(@NotNull PyExpression argument) {
return argument instanceof PyStarArgument && !((PyStarArgument)argument).isKeyword();
}
@Nullable
private static <T> T next(@NotNull List<T> list) {
return list.isEmpty() ? null : list.remove(0);
}
@NotNull
private static List<PyParameter> dropImplicitParameters(@NotNull List<PyParameter> parameters, int offset) {
final ArrayList<PyParameter> results = new ArrayList<PyParameter>(parameters);
for (int i = 0; i < offset && !results.isEmpty(); i++) {
results.remove(0);
}
return results;
}
/**
* Maps arguments of a call to parameters of a callee.
* must contain already resolved callee with flags set appropriately.
@@ -269,8 +269,8 @@ public class PyNamedParameterImpl extends PyBaseElementImpl<PyNamedParameterStub
final PyResolveContext resolveContext = PyResolveContext.noImplicits().withTypeEvalContext(context);
final PyArgumentList argumentList = call.getArgumentList();
if (argumentList != null) {
final CallArgumentsMapping mapping = argumentList.analyzeCall(resolveContext);
for (Map.Entry<PyExpression, PyNamedParameter> entry : mapping.getPlainMappedParams().entrySet()) {
final Map<PyExpression, PyNamedParameter> mapping = CallArgumentsMappingImpl.map(call, resolveContext);
for (Map.Entry<PyExpression, PyNamedParameter> entry : mapping.entrySet()) {
if (entry.getValue() == PyNamedParameterImpl.this) {
final PyExpression argument = entry.getKey();
if (argument != null) {
@@ -393,8 +393,8 @@ public class PyNamedParameterImpl extends PyBaseElementImpl<PyNamedParameterStub
}
}
final PyResolveContext resolveContext = PyResolveContext.noImplicits().withTypeEvalContext(context);
final CallArgumentsMapping mapping = argumentList.analyzeCall(resolveContext);
for (Map.Entry<PyExpression, PyNamedParameter> entry : mapping.getPlainMappedParams().entrySet()) {
final Map<PyExpression, PyNamedParameter> mapping = CallArgumentsMappingImpl.map(callExpression, resolveContext);
for (Map.Entry<PyExpression, PyNamedParameter> entry : mapping.entrySet()) {
if (entry.getKey() == element) {
return entry.getValue();
}
@@ -24,6 +24,7 @@ import com.intellij.util.ArrayUtil;
import com.jetbrains.python.PyNames;
import com.jetbrains.python.codeInsight.PyCustomMember;
import com.jetbrains.python.psi.*;
import com.jetbrains.python.psi.impl.CallArgumentsMappingImpl;
import com.jetbrains.python.psi.impl.PyBuiltinCache;
import com.jetbrains.python.psi.resolve.PyResolveContext;
import com.jetbrains.python.psi.resolve.RatedResolveResult;
@@ -456,9 +457,9 @@ public class PyTypeChecker {
final PyExpression callee = call.getCallee();
final PyArgumentList args = call.getArgumentList();
if (args != null) {
final CallArgumentsMapping mapping = args.analyzeCall(PyResolveContext.noImplicits().withTypeEvalContext(context));
final Map<PyExpression, PyNamedParameter> arguments = mapping.getPlainMappedParams();
final PyCallExpression.PyMarkedCallee markedCallee = mapping.getMarkedCallee();
final PyResolveContext resolveContext = PyResolveContext.noImplicits().withTypeEvalContext(context);
final Map<PyExpression, PyNamedParameter> arguments = CallArgumentsMappingImpl.map(call, resolveContext);
final PyCallExpression.PyMarkedCallee markedCallee = call.resolveCallee(resolveContext);
if (markedCallee != null) {
final PyCallable callable = markedCallee.getCallable();
if (callable instanceof PyFunction) {
@@ -47,6 +47,7 @@ import com.jetbrains.python.PyTokenTypes;
import com.jetbrains.python.PythonStringUtil;
import com.jetbrains.python.codeInsight.dataflow.scope.ScopeUtil;
import com.jetbrains.python.psi.*;
import com.jetbrains.python.psi.impl.CallArgumentsMappingImpl;
import com.jetbrains.python.psi.resolve.PyResolveContext;
import com.jetbrains.python.psi.types.PyNoneType;
import com.jetbrains.python.psi.types.PyType;
@@ -225,11 +226,16 @@ abstract public class IntroduceHandler implements RefactoringActionHandler {
final PyArgumentList argList = PsiTreeUtil.getParentOfType(expression, PyArgumentList.class);
if (argList != null) {
final CallArgumentsMapping result = argList.analyzeCall(PyResolveContext.noImplicits());
if (result.getMarkedCallee() != null) {
final PyNamedParameter namedParameter = result.getPlainMappedParams().get(expression);
if (namedParameter != null) {
candidates.add(namedParameter.getName());
final PyCallExpression callExpr = argList.getCallExpression();
if (callExpr != null) {
final PyResolveContext resolveContext = PyResolveContext.noImplicits();
final PyCallExpression.PyMarkedCallee markedCallee = callExpr.resolveCallee(resolveContext);
if (markedCallee != null) {
final Map<PyExpression, PyNamedParameter> mapping = CallArgumentsMappingImpl.map(callExpr, resolveContext);
final PyNamedParameter namedParameter = mapping.get(expression);
if (namedParameter != null) {
candidates.add(namedParameter.getName());
}
}
}
}
@@ -8,4 +8,4 @@ def f(spam, eggs):
def test():
f(<warning descr="Expected type 'list[Union[str, unicode]]', got 'list[int]' instead">[1, 2, 3]</warning>,
(<warning descr="Expected type 'Tuple[bool, int, unicode]', got 'Tuple[bool, int, str]' instead">False, 2, ''</warning>))
<warning descr="Expected type 'Tuple[bool, int, unicode]', got 'Tuple[bool, int, str]' instead">(False, 2, '')</warning>)