From 891c4bf405dad165a6d8c9daa0f15fcb72467107 Mon Sep 17 00:00:00 2001 From: Mikhail Golubev Date: Mon, 6 Nov 2017 20:10:55 +0300 Subject: [PATCH] PY-12002 KeywordArgumentCompletionUtil operates on type level --- .../KeywordArgumentCompletionUtil.java | 188 ++++++------------ .../python/psi/types/PyClassTypeImpl.java | 35 ++++ ...ameterWithDefaultAfterKeywordContainer2.py | 2 + .../MapReturnElementType.py | 2 +- 4 files changed, 98 insertions(+), 129 deletions(-) diff --git a/python/src/com/jetbrains/python/psi/impl/references/KeywordArgumentCompletionUtil.java b/python/src/com/jetbrains/python/psi/impl/references/KeywordArgumentCompletionUtil.java index 1f9e423c7d81..effcd3d75c43 100644 --- a/python/src/com/jetbrains/python/psi/impl/references/KeywordArgumentCompletionUtil.java +++ b/python/src/com/jetbrains/python/psi/impl/references/KeywordArgumentCompletionUtil.java @@ -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 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 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 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 collectParameterNamesFromType(@NotNull PyCallableType type, + @NotNull PyCallExpression callSite, + @NotNull TypeEvalContext context) { + List result = new ArrayList<>(); + if (type.isCallable()) { + final List 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 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 callableParameters = ((PyCallableType)memberType).getParameters(context); - if (callableParameters != null) { - fetchCallablesFromCallableType(ret, callExpr, callableParameters); - } - } - } - } - - private static void fetchCallablesFromCallableType(@NotNull final List ret, - @NotNull final PyCallExpression callExpr, - @NotNull final Iterable callableParameters) { - final List 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 ret) { - addKeywordArgumentVariants(callable, callExpr, ret, new HashSet<>()); - } - - public static void addKeywordArgumentVariants(PyCallable callable, PyCallExpression callExpr, List ret, - Collection visited) { - if (visited.contains(callable)) { - return; - } - visited.add(callable); - - final TypeEvalContext context = TypeEvalContext.codeCompletion(callable.getProject(), callable.getContainingFile()); - final List parameters = callable.getParameters(context); - - if (callable instanceof PyFunction) { - addKeywordArgumentVariantsForFunction(callExpr, ret, visited, (PyFunction)callable, parameters, context); - } - else { - final Collection 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 ret, - @NotNull final Collection parameterNames) { - for (final String parameterName : parameterNames) { - ret.add(PyUtil.createNamedParameterLookup(parameterName, callExpr.getProject())); - } - } - private static void addKeywordArgumentVariantsForFunction(@NotNull final PyCallExpression callExpr, - @NotNull final List ret, - @NotNull final Collection visited, @NotNull final PyFunction function, @NotNull final List parameters, + @NotNull final List ret, + @NotNull final Set 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 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 myRet; + private final List myRet; private boolean myHasSelf = false; private boolean myHasKwArgs = false; private PyParameter kwArgsParam = null; - public KwArgParameterCollector(boolean needSelf, List ret) { + public KwArgParameterCollector(boolean needSelf, List 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 myRet; + private final List myRet; private final PyParameter myKwArgs; private boolean kwArgsTransit = true; - public KwArgFromStatementCallCollector(List ret, @NotNull PyParameter kwArgs) { + public KwArgFromStatementCallCollector(List 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); } } } diff --git a/python/src/com/jetbrains/python/psi/types/PyClassTypeImpl.java b/python/src/com/jetbrains/python/psi/types/PyClassTypeImpl.java index 6daf69fd0ed1..477824c31d5d 100644 --- a/python/src/com/jetbrains/python/psi/types/PyClassTypeImpl.java +++ b/python/src/com/jetbrains/python/psi/types/PyClassTypeImpl.java @@ -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 getParameters(@NotNull TypeEvalContext context) { + if (isDefinition()) { + List 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 getParametersOfMethod(@NotNull String name, @NotNull TypeEvalContext context) { + final List 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()); diff --git a/python/testData/inspections/PyArgumentListInspection/parameterWithDefaultAfterKeywordContainer2.py b/python/testData/inspections/PyArgumentListInspection/parameterWithDefaultAfterKeywordContainer2.py index 2703fa92328b..8872a9e99666 100644 --- a/python/testData/inspections/PyArgumentListInspection/parameterWithDefaultAfterKeywordContainer2.py +++ b/python/testData/inspections/PyArgumentListInspection/parameterWithDefaultAfterKeywordContainer2.py @@ -2,6 +2,8 @@ kwargs = {'foo': 'bar'} class Foo(object): + def __init__(self, **kwargs): + pass @classmethod def test(cls): diff --git a/python/testData/inspections/PyTypeCheckerInspection/MapReturnElementType.py b/python/testData/inspections/PyTypeCheckerInspection/MapReturnElementType.py index 49c5527aa746..906357ede012 100644 --- a/python/testData/inspections/PyTypeCheckerInspection/MapReturnElementType.py +++ b/python/testData/inspections/PyTypeCheckerInspection/MapReturnElementType.py @@ -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 + ys[0], 'bar' + ys[1])