Integrate PyCallExpressionHelper.forEveryScopeTakeOverloadsOtherwiseImplementations into PyCallExpressionHelper.multiResolveRatedCallee

This commit is contained in:
Semyon Proshev
2017-05-13 00:17:49 +03:00
parent 4b7f436e13
commit eccfd264e3
4 changed files with 35 additions and 41 deletions
@@ -22,7 +22,6 @@ import com.intellij.openapi.util.Pair;
import com.intellij.openapi.util.TextRange;
import com.intellij.psi.PsiElement;
import com.intellij.psi.PsiFile;
import com.intellij.psi.ResolveResult;
import com.intellij.psi.util.PsiTreeUtil;
import com.intellij.util.ArrayUtil;
import com.intellij.util.text.CharArrayUtil;
@@ -76,10 +75,12 @@ public class PyParameterInfoHandler implements ParameterInfoHandler<PyArgumentLi
final List<PyCallExpression.PyRatedMarkedCallee> ratedMarkedCallees =
PyUtil.filterTopPriorityResults(call.multiResolveRatedCallee(resolveContext));
final Object[] items = PyCallExpressionHelper
.forEveryScopeTakeOverloadsOtherwiseImplementations(ratedMarkedCallees, ResolveResult::getElement, typeEvalContext)
.map(ratedMarkedCallee -> Pair.createNonNull(call, ratedMarkedCallee.getMarkedCallee()))
.toArray();
final Object[] items = new Object[ratedMarkedCallees.size()];
int currentPosition = 0;
for (PyCallExpression.PyRatedMarkedCallee ratedMarkedCallee : ratedMarkedCallees) {
items[currentPosition] = Pair.createNonNull(call, ratedMarkedCallee.getMarkedCallee());
currentPosition++;
}
context.setItemsToShow(items);
@@ -31,7 +31,6 @@ import com.jetbrains.python.PyTokenTypes;
import com.jetbrains.python.inspections.quickfix.PyRemoveArgumentQuickFix;
import com.jetbrains.python.inspections.quickfix.PyRenameArgumentQuickFix;
import com.jetbrains.python.psi.*;
import com.jetbrains.python.psi.impl.PyCallExpressionHelper;
import com.jetbrains.python.psi.resolve.PyResolveContext;
import com.jetbrains.python.psi.types.PyABCUtil;
import com.jetbrains.python.psi.types.PyType;
@@ -120,7 +119,8 @@ public class PyArgumentListInspection extends PyInspection {
final PyCallExpression call = node.getCallExpression();
if (call == null) return;
final List<PyCallExpression.PyArgumentsMapping> mappings = calculateMappings(call, context, implicitOffset);
final PyResolveContext resolveContext = PyResolveContext.noImplicits().withTypeEvalContext(context);
final List<PyCallExpression.PyArgumentsMapping> mappings = call.multiMapArguments(resolveContext, implicitOffset);
for (PyCallExpression.PyArgumentsMapping mapping : mappings) {
final PyCallExpression.PyMarkedCallee callee = mapping.getMarkedCallee();
@@ -146,20 +146,6 @@ public class PyArgumentListInspection extends PyInspection {
inspectPyArgumentList(node, holder, context, 0);
}
@NotNull
private static List<PyCallExpression.PyArgumentsMapping> calculateMappings(@NotNull PyCallExpression call,
@NotNull TypeEvalContext context,
int implicitOffset) {
final PyResolveContext resolveContext = PyResolveContext.noImplicits().withTypeEvalContext(context);
final List<PyCallExpression.PyRatedMarkedCallee> ratedMarkedCallees =
PyUtil.filterTopPriorityResults(call.multiResolveRatedCallee(resolveContext, implicitOffset));
return PyCallExpressionHelper
.forEveryScopeTakeOverloadsOtherwiseImplementations(ratedMarkedCallees, ResolveResult::getElement, context)
.map(ratedMarkedCallee -> PyCallExpressionHelper.mapArguments(call, ratedMarkedCallee.getMarkedCallee(), context))
.collect(Collectors.toList());
}
private static boolean decoratedClassInitCall(@Nullable PyExpression callee, @NotNull PyFunction function) {
if (callee instanceof PyReferenceExpression && PyUtil.isInit(function)) {
final PsiPolyVariantReference classReference = ((PyReferenceExpression)callee).getReference();
@@ -119,19 +119,29 @@ public class PyCallExpressionHelper {
public static List<PyCallExpression.PyRatedMarkedCallee> multiResolveRatedCallee(@NotNull PyCallExpression call,
@NotNull PyResolveContext resolveContext,
int implicitOffset) {
return multiResolveCallee(call.getCallee(), resolveContext)
.stream()
.map(resolveResult -> clarifyResolveResult(resolveResult, resolveContext, call))
.filter(Objects::nonNull)
.map(resolveResult -> markResolveResult(resolveResult, resolveContext.getTypeEvalContext(), implicitOffset))
.filter(Objects::nonNull)
final LinkedHashMap<PsiElement, PyCallExpression.PyRatedMarkedCallee> result = new LinkedHashMap<>();
for (QualifiedRatedResolveResult resolveResult : multiResolveCallee(call.getCallee(), resolveContext)) {
final ClarifiedResolveResult clarifiedResolveResult = clarifyResolveResult(resolveResult, resolveContext, call);
if (clarifiedResolveResult == null) continue;
final PyCallExpression.PyRatedMarkedCallee markedCallee =
markResolveResult(clarifiedResolveResult, resolveContext.getTypeEvalContext(), implicitOffset);
if (markedCallee == null) continue;
// while clarifying resolve results we could get duplicate callables so we have to group them and select result with highest rate
.collect(Collectors.groupingBy(markedCallee -> markedCallee.getElement(), LinkedHashMap::new, Collectors.toList()))
.entrySet()
.stream()
.map(entry -> entry.getValue().stream().max(Comparator.comparingInt(PyCallExpression.PyRatedMarkedCallee::getRate)).orElse(null))
.filter(Objects::nonNull)
.collect(Collectors.toList());
result.put(markedCallee.getElement(),
StreamEx
.of(markedCallee, result.get(markedCallee.getElement()))
.nonNull()
.max(Comparator.comparingInt(PyCallExpression.PyRatedMarkedCallee::getRate))
.orElse(null));
}
return StreamEx
.of(forEveryScopeTakeOverloadsOtherwiseImplementations(result.entrySet(), Map.Entry::getKey, resolveContext.getTypeEvalContext()))
.map(Map.Entry::getValue)
.toList();
}
@NotNull
@@ -842,9 +852,9 @@ public class PyCallExpressionHelper {
}
@NotNull
public static <E> Stream<E> forEveryScopeTakeOverloadsOtherwiseImplementations(@NotNull List<E> elements,
@NotNull Function<? super E, PsiElement> mapper,
@NotNull TypeEvalContext context) {
private static <E> Stream<E> forEveryScopeTakeOverloadsOtherwiseImplementations(@NotNull Collection<E> elements,
@NotNull Function<? super E, PsiElement> mapper,
@NotNull TypeEvalContext context) {
if (!containsOverloadsAndImplementations(elements, mapper, context)) {
return elements.stream();
}
@@ -857,7 +867,7 @@ public class PyCallExpressionHelper {
.flatMap(oneScopeElements -> takeOverloadsOtherwiseImplementations(oneScopeElements, mapper, context));
}
private static <E> boolean containsOverloadsAndImplementations(@NotNull List<E> elements,
private static <E> boolean containsOverloadsAndImplementations(@NotNull Collection<E> elements,
@NotNull Function<? super E, PsiElement> mapper,
@NotNull TypeEvalContext context) {
boolean containsOverloads = false;
@@ -18,7 +18,6 @@ package com.jetbrains.python.psi.types;
import com.intellij.openapi.extensions.Extensions;
import com.intellij.psi.PsiElement;
import com.intellij.psi.PsiNamedElement;
import com.intellij.psi.ResolveResult;
import com.intellij.util.ArrayUtil;
import com.intellij.util.containers.ContainerUtil;
import com.jetbrains.python.PyNames;
@@ -628,9 +627,7 @@ public class PyTypeChecker {
final List<PyCallExpression.PyRatedMarkedCallee> ratedMarkedCallees =
PyUtil.filterTopPriorityResults(((PyCallExpression)callSite).multiResolveRatedCallee(resolveContext));
return forEveryScopeTakeOverloadsOtherwiseImplementations(ratedMarkedCallees, ResolveResult::getElement, context)
.map(PyCallExpression.PyRatedMarkedCallee::getElement)
.collect(Collectors.toList());
return ContainerUtil.map(ratedMarkedCallees, PyCallExpression.PyRatedMarkedCallee::getElement);
}
else if (callSite instanceof PySubscriptionExpression || callSite instanceof PyBinaryExpression) {
final List<PyCallable> results = new ArrayList<>();