PY-12002 KeywordArgumentCompletionUtil operates on type level

This commit is contained in:
Mikhail Golubev
2017-12-01 17:46:57 +03:00
parent cc97c9fc4e
commit 891c4bf405
4 changed files with 98 additions and 129 deletions
@@ -20,8 +20,8 @@ import com.intellij.openapi.extensions.Extensions;
import com.intellij.openapi.util.Comparing;
import com.intellij.psi.PsiElement;
import com.intellij.psi.util.PsiTreeUtil;
import com.intellij.util.containers.ContainerUtil;
import com.jetbrains.python.PyNames;
import com.jetbrains.python.codeInsight.stdlib.PyNamedTupleType;
import com.jetbrains.python.psi.*;
import com.jetbrains.python.psi.impl.PyKeywordArgumentProvider;
import com.jetbrains.python.psi.resolve.PyResolveContext;
@@ -30,12 +30,13 @@ import com.jetbrains.python.psi.search.PySuperMethodsSearch;
import com.jetbrains.python.psi.types.*;
import one.util.streamex.StreamEx;
import org.jetbrains.annotations.NotNull;
import org.jetbrains.annotations.Nullable;
import java.util.ArrayList;
import java.util.Collection;
import java.util.HashSet;
import java.util.List;
import java.util.Set;
import static com.jetbrains.python.psi.PyUtil.as;
public class KeywordArgumentCompletionUtil {
public static void collectFunctionArgNames(PyElement element, List<LookupElement> ret, @NotNull final TypeEvalContext context) {
@@ -43,137 +44,72 @@ public class KeywordArgumentCompletionUtil {
if (callExpr != null) {
PyExpression callee = callExpr.getCallee();
if (callee instanceof PyReferenceExpression && element.getParent() == callExpr.getArgumentList()) {
PsiElement def = getElementByType(context, callee);
if (def == null) {
def = getElementByChain(context, (PyReferenceExpression)callee);
}
if (def instanceof PyCallable) {
addKeywordArgumentVariants((PyCallable)def, callExpr, ret);
}
else if (def instanceof PyClass) {
PyFunction init = ((PyClass)def).findMethodByName(PyNames.INIT, true, null); // search in superclasses
if (init != null) {
addKeywordArgumentVariants(init, callExpr, ret);
PyType calleeType = context.getType(callee);
if (calleeType == null) {
final PyTypedElement implicit = as(getElementByChain((PyReferenceExpression)callee, context), PyTypedElement.class);
if (implicit != null) {
calleeType = context.getType(implicit);
}
}
final PyType calleeType = context.getType(callee);
final PyUnionType unionType = PyUtil.as(calleeType, PyUnionType.class);
if (unionType != null) {
fetchCallablesFromUnion(ret, callExpr, unionType, context);
final StreamEx<PyType> types;
if (calleeType instanceof PyUnionType) {
types = StreamEx.of(((PyUnionType)calleeType).getMembers());
}
final PyNamedTupleType namedTupleType = PyUtil.as(calleeType, PyNamedTupleType.class);
if (namedTupleType != null) {
for (String name : namedTupleType.getFields().keySet()) {
ret.add(
PyUtil.createNamedParameterLookup(name, element.getProject())
);
}
else {
types = StreamEx.of(calleeType);
}
final List<LookupElement> extra = types
.select(PyCallableType.class)
.flatMap(type -> collectParameterNamesFromType(type, callExpr, context).stream())
.map(name -> PyUtil.createNamedParameterLookup(name, element.getProject()))
.toList();
ret.addAll(extra);
}
}
}
@Nullable
private static PyElement getElementByType(@NotNull final TypeEvalContext context, @NotNull final PyExpression callee) {
final PyType pyType = context.getType(callee);
if (pyType instanceof PyFunctionType) {
return ((PyFunctionType)pyType).getCallable();
@NotNull
private static List<String> collectParameterNamesFromType(@NotNull PyCallableType type,
@NotNull PyCallExpression callSite,
@NotNull TypeEvalContext context) {
List<String> result = new ArrayList<>();
if (type.isCallable()) {
final List<PyCallableParameter> parameters = type.getParameters(context);
if (parameters != null) {
for (PyCallableParameter parameter : parameters) {
if (parameter.isKeywordContainer() || parameter.isPositionalContainer()) {
continue;
}
ContainerUtil.addIfNotNull(result, parameter.getName());
}
if (type instanceof PyFunctionType) {
final PyFunction func = as(((PyFunctionType)type).getCallable(), PyFunction.class);
if (func != null) {
addKeywordArgumentVariantsForFunction(callSite, func, parameters, result, new HashSet<>(), context);
}
}
}
}
if (pyType instanceof PyClassType) {
return ((PyClassType)pyType).getPyClass();
}
return null;
return result;
}
private static PsiElement getElementByChain(@NotNull TypeEvalContext context, PyReferenceExpression callee) {
private static PsiElement getElementByChain(@NotNull PyReferenceExpression callee, @NotNull TypeEvalContext context) {
final PyResolveContext resolveContext = PyResolveContext.defaultContext().withTypeEvalContext(context);
final QualifiedResolveResult result = callee.followAssignmentsChain(resolveContext);
return result.getElement();
}
private static void fetchCallablesFromUnion(@NotNull final List<LookupElement> ret,
@NotNull final PyCallExpression callExpr,
@NotNull final PyUnionType unionType,
@NotNull final TypeEvalContext context) {
for (final PyType memberType : unionType.getMembers()) {
if (memberType instanceof PyUnionType) {
fetchCallablesFromUnion(ret, callExpr, (PyUnionType)memberType, context);
}
if (memberType instanceof PyFunctionType) {
final PyFunctionType type = (PyFunctionType)memberType;
if (type.isCallable()) {
addKeywordArgumentVariants(type.getCallable(), callExpr, ret);
}
}
if (memberType instanceof PyCallableType) {
final List<PyCallableParameter> callableParameters = ((PyCallableType)memberType).getParameters(context);
if (callableParameters != null) {
fetchCallablesFromCallableType(ret, callExpr, callableParameters);
}
}
}
}
private static void fetchCallablesFromCallableType(@NotNull final List<LookupElement> ret,
@NotNull final PyCallExpression callExpr,
@NotNull final Iterable<PyCallableParameter> callableParameters) {
final List<String> parameterNames = new ArrayList<>();
for (final PyCallableParameter callableParameter : callableParameters) {
final String name = callableParameter.getName();
if (name != null) {
parameterNames.add(name);
}
}
addKeywordArgumentVariantsForCallable(callExpr, ret, parameterNames);
}
public static void addKeywordArgumentVariants(PyCallable callable, PyCallExpression callExpr, final List<LookupElement> ret) {
addKeywordArgumentVariants(callable, callExpr, ret, new HashSet<>());
}
public static void addKeywordArgumentVariants(PyCallable callable, PyCallExpression callExpr, List<LookupElement> ret,
Collection<PyCallable> visited) {
if (visited.contains(callable)) {
return;
}
visited.add(callable);
final TypeEvalContext context = TypeEvalContext.codeCompletion(callable.getProject(), callable.getContainingFile());
final List<PyCallableParameter> parameters = callable.getParameters(context);
if (callable instanceof PyFunction) {
addKeywordArgumentVariantsForFunction(callExpr, ret, visited, (PyFunction)callable, parameters, context);
}
else {
final Collection<String> parameterNames = new ArrayList<>();
for (final PyCallableParameter parameter : parameters) {
final String name = parameter.getName();
if (name != null) {
parameterNames.add(name);
}
}
addKeywordArgumentVariantsForCallable(callExpr, ret, parameterNames);
}
}
private static void addKeywordArgumentVariantsForCallable(@NotNull final PyCallExpression callExpr,
@NotNull final List<LookupElement> ret,
@NotNull final Collection<String> parameterNames) {
for (final String parameterName : parameterNames) {
ret.add(PyUtil.createNamedParameterLookup(parameterName, callExpr.getProject()));
}
}
private static void addKeywordArgumentVariantsForFunction(@NotNull final PyCallExpression callExpr,
@NotNull final List<LookupElement> ret,
@NotNull final Collection<PyCallable> visited,
@NotNull final PyFunction function,
@NotNull final List<PyCallableParameter> parameters,
@NotNull final List<String> ret,
@NotNull final Set<PyCallable> visited,
@NotNull final TypeEvalContext context) {
if (visited.contains(function)) {
return;
}
boolean needSelf = function.getContainingClass() != null && function.getModifier() != PyFunction.Modifier.STATICMETHOD;
final KwArgParameterCollector collector = new KwArgParameterCollector(needSelf, ret);
@@ -185,10 +121,7 @@ public class KeywordArgumentCompletionUtil {
if (collector.hasKwArgs()) {
for (PyKeywordArgumentProvider provider : Extensions.getExtensions(PyKeywordArgumentProvider.EP_NAME)) {
final List<String> arguments = provider.getKeywordArguments(function, callExpr);
for (String argument : arguments) {
ret.add(PyUtil.createNamedParameterLookup(argument, callExpr.getProject()));
}
ret.addAll(provider.getKeywordArguments(function, callExpr));
}
KwArgFromStatementCallCollector fromStatementCallCollector = new KwArgFromStatementCallCollector(ret, collector.getKwArgs());
function.getStatementList().acceptChildren(fromStatementCallCollector);
@@ -197,9 +130,9 @@ public class KeywordArgumentCompletionUtil {
// nothing interesting besides self and **kwargs, let's look at superclass (PY-778)
if (fromStatementCallCollector.isKwArgsTransit()) {
final PsiElement superMethod = PySuperMethodsSearch.search(function, context).findFirst();
if (superMethod instanceof PyFunction) {
addKeywordArgumentVariants((PyFunction)superMethod, callExpr, ret, visited);
final PyFunction superMethod = as(PySuperMethodsSearch.search(function, context).findFirst(), PyFunction.class);
if (superMethod != null) {
addKeywordArgumentVariantsForFunction(callExpr, superMethod, superMethod.getParameters(context), ret, visited, context);
}
}
}
@@ -208,12 +141,12 @@ public class KeywordArgumentCompletionUtil {
public static class KwArgParameterCollector extends PyElementVisitor {
private int myCount;
private final boolean myNeedSelf;
private final List<LookupElement> myRet;
private final List<String> myRet;
private boolean myHasSelf = false;
private boolean myHasKwArgs = false;
private PyParameter kwArgsParam = null;
public KwArgParameterCollector(boolean needSelf, List<LookupElement> ret) {
public KwArgParameterCollector(boolean needSelf, List<String> ret) {
myNeedSelf = needSelf;
myRet = ret;
}
@@ -228,8 +161,7 @@ public class KeywordArgumentCompletionUtil {
PyNamedParameter namedParam = par.getAsNamed();
if (namedParam != null) {
if (!namedParam.isKeywordContainer() && !namedParam.isPositionalContainer()) {
final LookupElement item = PyUtil.createNamedParameterLookup(namedParam.getName(), par.getProject());
myRet.add(item);
myRet.add(namedParam.getName());
}
else if (namedParam.isKeywordContainer()) {
myHasKwArgs = true;
@@ -255,11 +187,11 @@ public class KeywordArgumentCompletionUtil {
}
public static class KwArgFromStatementCallCollector extends PyElementVisitor {
private final List<LookupElement> myRet;
private final List<String> myRet;
private final PyParameter myKwArgs;
private boolean kwArgsTransit = true;
public KwArgFromStatementCallCollector(List<LookupElement> ret, @NotNull PyParameter kwArgs) {
public KwArgFromStatementCallCollector(List<String> ret, @NotNull PyParameter kwArgs) {
myRet = ret;
this.myKwArgs = kwArgs;
}
@@ -307,7 +239,7 @@ public class KeywordArgumentCompletionUtil {
argument instanceof PyStringLiteralExpression) {
String name = ((PyStringLiteralExpression)argument).getStringValue();
if (PyNames.isIdentifier(name)) {
myRet.add(PyUtil.createNamedParameterLookup(name, argument.getProject()));
myRet.add(name);
}
}
}
@@ -31,6 +31,7 @@ import com.jetbrains.python.psi.resolve.PyResolveContext;
import com.jetbrains.python.psi.resolve.PyResolveProcessor;
import com.jetbrains.python.psi.resolve.RatedResolveResult;
import com.jetbrains.python.toolbox.Maybe;
import one.util.streamex.StreamEx;
import org.jetbrains.annotations.NotNull;
import org.jetbrains.annotations.Nullable;
@@ -461,6 +462,40 @@ public class PyClassTypeImpl extends UserDataHolderBase implements PyClassType {
return false;
}
@Nullable
@Override
public List<PyCallableParameter> getParameters(@NotNull TypeEvalContext context) {
if (isDefinition()) {
List<PyCallableParameter> params = getParametersOfMethod(PyNames.INIT, context);
if (params == null) {
// TODO better way to resolve the constructor method here
params = getParametersOfMethod(PyNames.NEW, context);
}
if (params != null) {
// Skip "self" for __init__ and "cls" for __new__
return params.subList(1, params.size());
}
return null;
}
return getParametersOfMethod(PyNames.CALL, context);
}
@Nullable
private List<PyCallableParameter> getParametersOfMethod(@NotNull String name, @NotNull TypeEvalContext context) {
final List<? extends RatedResolveResult> results =
resolveMember(name, null, AccessDirection.READ, PyResolveContext.noImplicits().withTypeEvalContext(context), true);
if (results != null) {
return StreamEx.of(results)
.map(RatedResolveResult::getElement)
.select(PyCallable.class)
.map(func -> func.getParameters(context))
.findFirst()
.orElse(null);
}
return null;
}
private static boolean isMethodType(@NotNull PyClassType type) {
final PyBuiltinCache builtinCache = PyBuiltinCache.getInstance(type.getPyClass());
return type.equals(builtinCache.getClassMethodType()) || type.equals(builtinCache.getStaticMethodType());
@@ -2,6 +2,8 @@ kwargs = {'foo': 'bar'}
class Foo(object):
def __init__(self, **kwargs):
pass
@classmethod
def test(cls):
@@ -1,5 +1,5 @@
def test():
xs = map(lambda x: x + 1, [1, 2, 3])
print('foo' + xs[0])
ys = map(tuple, iter([1, 2, 3]))
ys = map(tuple, iter([[1, 2, 3]]))
print(1 + <warning descr="Expected type 'int', got 'tuple' instead">ys[0]</warning>, 'bar' + <warning descr="Expected type 'AnyStr', got 'tuple' instead">ys[1]</warning>)