diff --git a/python/src/com/jetbrains/python/inspections/PyArgumentListInspection.java b/python/src/com/jetbrains/python/inspections/PyArgumentListInspection.java index ca418855fbb5..345bcef9938f 100644 --- a/python/src/com/jetbrains/python/inspections/PyArgumentListInspection.java +++ b/python/src/com/jetbrains/python/inspections/PyArgumentListInspection.java @@ -31,6 +31,7 @@ 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; @@ -45,6 +46,7 @@ import java.util.*; import java.util.stream.Collectors; public class PyArgumentListInspection extends PyInspection { + @Override @Nls @NotNull public String getDisplayName() { @@ -118,8 +120,7 @@ public class PyArgumentListInspection extends PyInspection { final PyCallExpression call = node.getCallExpression(); if (call == null) return; - final PyResolveContext resolveContext = PyResolveContext.noImplicits().withTypeEvalContext(context); - final List mappings = call.multiMapArguments(resolveContext, implicitOffset); + final List mappings = calculateMappings(call, context, implicitOffset); for (PyCallExpression.PyArgumentsMapping mapping : mappings) { final PyCallExpression.PyMarkedCallee callee = mapping.getMarkedCallee(); @@ -145,6 +146,20 @@ 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, 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(); @@ -268,11 +283,10 @@ public class PyArgumentListInspection extends PyInspection { .map(PyCallExpression.PyArgumentsMapping::getMarkedCallee) .nonNull() .map(markedCallee -> calculatePossibleCalleeRepresentation(markedCallee.getCallable(), context)) - .nonNull() .collect(Collectors.joining("
")); } - @Nullable + @NotNull private static String calculatePossibleCalleeRepresentation(@NotNull PyCallable callable, @NotNull TypeEvalContext context) { final String name = callable.getName(); final String parameters = callable.getParameterList().getPresentableText(true, context); diff --git a/python/testData/inspections/PyArgumentListInspection/OverloadsAndImplementationInImportedClass/b.py b/python/testData/inspections/PyArgumentListInspection/OverloadsAndImplementationInImportedClass/b.py new file mode 100644 index 000000000000..97dfc5be4c61 --- /dev/null +++ b/python/testData/inspections/PyArgumentListInspection/OverloadsAndImplementationInImportedClass/b.py @@ -0,0 +1,3 @@ +import c + +c.A().foo() \ No newline at end of file diff --git a/python/testData/inspections/PyArgumentListInspection/OverloadsAndImplementationInImportedClass/c.py b/python/testData/inspections/PyArgumentListInspection/OverloadsAndImplementationInImportedClass/c.py new file mode 100644 index 000000000000..ef6d14f66b95 --- /dev/null +++ b/python/testData/inspections/PyArgumentListInspection/OverloadsAndImplementationInImportedClass/c.py @@ -0,0 +1,18 @@ +from typing import overload + + +class A: + @overload + def foo(self, value: None) -> None: + pass + + @overload + def foo(self, value: int) -> str: + pass + + @overload + def foo(self, value: str) -> str: + pass + + def foo(self, value): + return None \ No newline at end of file diff --git a/python/testData/inspections/PyArgumentListInspection/OverloadsAndImplementationInImportedModule/b.py b/python/testData/inspections/PyArgumentListInspection/OverloadsAndImplementationInImportedModule/b.py new file mode 100644 index 000000000000..1c9cc9f204de --- /dev/null +++ b/python/testData/inspections/PyArgumentListInspection/OverloadsAndImplementationInImportedModule/b.py @@ -0,0 +1,3 @@ +import c + +c.foo() \ No newline at end of file diff --git a/python/testData/inspections/PyArgumentListInspection/OverloadsAndImplementationInImportedModule/c.py b/python/testData/inspections/PyArgumentListInspection/OverloadsAndImplementationInImportedModule/c.py new file mode 100644 index 000000000000..6fcd2dd29a54 --- /dev/null +++ b/python/testData/inspections/PyArgumentListInspection/OverloadsAndImplementationInImportedModule/c.py @@ -0,0 +1,17 @@ +from typing import overload + + +@overload +def foo(value: None) -> None: + pass + +@overload +def foo(value: int) -> str: + pass + +@overload +def foo(value: str) -> str: + pass + +def foo(value): + return None \ No newline at end of file diff --git a/python/testData/inspections/PyArgumentListInspection/overloadsAndImplementationInClass.py b/python/testData/inspections/PyArgumentListInspection/overloadsAndImplementationInClass.py new file mode 100644 index 000000000000..4ba41c7f33ce --- /dev/null +++ b/python/testData/inspections/PyArgumentListInspection/overloadsAndImplementationInClass.py @@ -0,0 +1,21 @@ +from typing import overload + + +class A: + @overload + def foo(self, value: None) -> None: + pass + + @overload + def foo(self, value: int) -> str: + pass + + @overload + def foo(self, value: str) -> str: + pass + + def foo(self, value): + return None + + +A().foo() \ No newline at end of file diff --git a/python/testData/inspections/PyArgumentListInspection/topLevelOverloadsAndImplementation.py b/python/testData/inspections/PyArgumentListInspection/topLevelOverloadsAndImplementation.py new file mode 100644 index 000000000000..dd9dbffb2c35 --- /dev/null +++ b/python/testData/inspections/PyArgumentListInspection/topLevelOverloadsAndImplementation.py @@ -0,0 +1,20 @@ +from typing import overload + + +@overload +def foo(value: None) -> None: + pass + +@overload +def foo(value: int) -> str: + pass + +@overload +def foo(value: str) -> str: + pass + +def foo(value): + return None + + +foo() \ No newline at end of file diff --git a/python/testSrc/com/jetbrains/python/inspections/PyArgumentListInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/PyArgumentListInspectionTest.java index 0d3b042e2b61..9e8ef918c13e 100644 --- a/python/testSrc/com/jetbrains/python/inspections/PyArgumentListInspectionTest.java +++ b/python/testSrc/com/jetbrains/python/inspections/PyArgumentListInspectionTest.java @@ -273,6 +273,26 @@ public class PyArgumentListInspectionTest extends PyTestCase { doTest(); } + // PY-22971 + public void testOverloadsAndImplementationInClass() { + runWithLanguageLevel(LanguageLevel.PYTHON35, this::doTest); + } + + // PY-22971 + public void testTopLevelOverloadsAndImplementation() { + runWithLanguageLevel(LanguageLevel.PYTHON35, this::doTest); + } + + // PY-22971 + public void testOverloadsAndImplementationInImportedClass() { + runWithLanguageLevel(LanguageLevel.PYTHON35, this::doMultiFileTest); + } + + // PY-22971 + public void testOverloadsAndImplementationInImportedModule() { + runWithLanguageLevel(LanguageLevel.PYTHON35, this::doMultiFileTest); + } + private void doMultiFileTest() { final String folderPath = "inspections/PyArgumentListInspection/" + getTestName(false) + "/";