diff --git a/python/src/com/jetbrains/python/PyParameterInfoHandler.java b/python/src/com/jetbrains/python/PyParameterInfoHandler.java index 5c4076b4b536..518051e25b80 100644 --- a/python/src/com/jetbrains/python/PyParameterInfoHandler.java +++ b/python/src/com/jetbrains/python/PyParameterInfoHandler.java @@ -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 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); diff --git a/python/src/com/jetbrains/python/inspections/PyArgumentListInspection.java b/python/src/com/jetbrains/python/inspections/PyArgumentListInspection.java index b148c78eb811..d3f80dc20f00 100644 --- a/python/src/com/jetbrains/python/inspections/PyArgumentListInspection.java +++ b/python/src/com/jetbrains/python/inspections/PyArgumentListInspection.java @@ -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 mappings = calculateMappings(call, context, implicitOffset); + final PyResolveContext resolveContext = PyResolveContext.noImplicits().withTypeEvalContext(context); + final List 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 calculateMappings(@NotNull PyCallExpression call, - @NotNull TypeEvalContext context, - int implicitOffset) { - final PyResolveContext resolveContext = PyResolveContext.noImplicits().withTypeEvalContext(context); - final List 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(); diff --git a/python/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.java b/python/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.java index 9fe41fb6a917..6577bc116282 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.java +++ b/python/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.java @@ -119,19 +119,29 @@ public class PyCallExpressionHelper { public static List 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 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 Stream forEveryScopeTakeOverloadsOtherwiseImplementations(@NotNull List elements, - @NotNull Function mapper, - @NotNull TypeEvalContext context) { + private static Stream forEveryScopeTakeOverloadsOtherwiseImplementations(@NotNull Collection elements, + @NotNull Function 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 boolean containsOverloadsAndImplementations(@NotNull List elements, + private static boolean containsOverloadsAndImplementations(@NotNull Collection elements, @NotNull Function mapper, @NotNull TypeEvalContext context) { boolean containsOverloads = false; diff --git a/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java b/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java index e165c112b574..93ab6cfc6365 100644 --- a/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java +++ b/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java @@ -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 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 results = new ArrayList<>();