diff --git a/python/src/com/jetbrains/python/codeInsight/override/PyOverrideImplementUtil.java b/python/src/com/jetbrains/python/codeInsight/override/PyOverrideImplementUtil.java index 57f07f6b793f..970d5e7c43df 100644 --- a/python/src/com/jetbrains/python/codeInsight/override/PyOverrideImplementUtil.java +++ b/python/src/com/jetbrains/python/codeInsight/override/PyOverrideImplementUtil.java @@ -169,14 +169,14 @@ public final class PyOverrideImplementUtil { final TypeEvalContext context = TypeEvalContext.userInitiated(cls.getProject(), cls.getContainingFile()); final PyFunction function = buildOverriddenFunction(cls, baseFunction, implement).addFunctionAfter(statementList, anchor); - addImports(baseFunction, function, context); + addImports(baseFunction, function); PyiUtil .getOverloads(baseFunction, context) .forEach( baseOverload -> { final PyFunction overload = (PyFunction)statementList.addBefore(baseOverload, function); - addImports(baseOverload, overload, context); + addImports(baseOverload, overload); } ); @@ -336,36 +336,20 @@ public final class PyOverrideImplementUtil { * @param baseFunction base function used to resolve types * @param function overridden function */ - private static void addImports(@NotNull PyFunction baseFunction, @NotNull PyFunction function, @NotNull TypeEvalContext context) { - final UnresolvedExpressionVisitor unresolvedExpressionVisitor = new UnresolvedExpressionVisitor(); - getAnnotations(function, context).forEach(annotation -> annotation.accept(unresolvedExpressionVisitor)); - getDecorators(function).forEach(decorator -> decorator.accept(unresolvedExpressionVisitor)); + private static void addImports(@NotNull PyFunction baseFunction, @NotNull PyFunction function) { + UnresolvedExpressionVisitor unresolvedExpressionVisitor = new UnresolvedExpressionVisitor() { + @Override + public void visitPyStatementList(@NotNull PyStatementList node) { + } + }; + function.accept(unresolvedExpressionVisitor); - final ResolveExpressionVisitor resolveExpressionVisitor = new ResolveExpressionVisitor(unresolvedExpressionVisitor.getUnresolved()); - getAnnotations(baseFunction, context).forEach(annotation -> annotation.accept(resolveExpressionVisitor)); - getDecorators(baseFunction).forEach(decorator -> decorator.accept(resolveExpressionVisitor)); - } - - /** - * Collect annotations from function parameters and return. - * - */ - @NotNull - private static List getAnnotations(@NotNull PyFunction function, @NotNull TypeEvalContext typeEvalContext) { - return StreamEx.of(function.getParameters(typeEvalContext)) - .map(PyCallableParameter::getParameter) - .select(PyNamedParameter.class) - .remove(PyParameter::isSelf) - .map(PyAnnotationOwner::getAnnotation) - .append(function.getAnnotation()) - .nonNull() - .toList(); - } - - @NotNull - private static List getDecorators(@NotNull PyFunction function) { - final PyDecoratorList decoratorList = function.getDecoratorList(); - return decoratorList == null ? Collections.emptyList() : Arrays.asList(decoratorList.getDecorators()); + ResolveExpressionVisitor resolveExpressionVisitor = new ResolveExpressionVisitor(unresolvedExpressionVisitor.getUnresolved()) { + @Override + public void visitPyStatementList(@NotNull PyStatementList node) { + } + }; + baseFunction.accept(resolveExpressionVisitor); } /** diff --git a/python/testData/override/ImportForParameterDefaultValue/main.py b/python/testData/override/ImportForParameterDefaultValue/main.py new file mode 100644 index 000000000000..6376f55f9e7c --- /dev/null +++ b/python/testData/override/ImportForParameterDefaultValue/main.py @@ -0,0 +1,5 @@ +from mod import Base + + +class Sub(Base): + pass diff --git a/python/testData/override/ImportForParameterDefaultValue/main_after.py b/python/testData/override/ImportForParameterDefaultValue/main_after.py new file mode 100644 index 000000000000..7c970ea5e67a --- /dev/null +++ b/python/testData/override/ImportForParameterDefaultValue/main_after.py @@ -0,0 +1,6 @@ +from mod import Base, default + + +class Sub(Base): + def method(self, param=default): + super().method(param) diff --git a/python/testData/override/ImportForParameterDefaultValue/mod.py b/python/testData/override/ImportForParameterDefaultValue/mod.py new file mode 100644 index 000000000000..32324bad0e29 --- /dev/null +++ b/python/testData/override/ImportForParameterDefaultValue/mod.py @@ -0,0 +1,6 @@ +default = object() + + +class Base: + def method(self, param=default): + pass diff --git a/python/testSrc/com/jetbrains/python/PyOverrideTest.java b/python/testSrc/com/jetbrains/python/PyOverrideTest.java index fc91b487d3bb..dc72a159771f 100644 --- a/python/testSrc/com/jetbrains/python/PyOverrideTest.java +++ b/python/testSrc/com/jetbrains/python/PyOverrideTest.java @@ -244,6 +244,13 @@ public class PyOverrideTest extends PyTestCase { }); } + public void testImportForParameterDefaultValue() { + myFixture.copyDirectoryToProject(getTestName(false), ""); + myFixture.configureByFile("main.py"); + doOverride(null); + myFixture.checkResultByFile(getTestName(false) + "/main_after.py", true); + } + @Override protected String getTestDataPath() { return super.getTestDataPath() + "/override";