PY-45588 Generalize PyOverrideImplementUtil.addImports to cover parameter defaults

GitOrigin-RevId: 828305e2a7449704c08503d8d68394d28ce98c83
This commit is contained in:
Mikhail Golubev
2023-08-09 20:53:35 +00:00
committed by intellij-monorepo-bot
parent b50c617b7c
commit c7bba04743
5 changed files with 39 additions and 31 deletions
@@ -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<PyAnnotation> 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<PyDecorator> 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);
}
/**
@@ -0,0 +1,5 @@
from mod import Base
class Sub(Base):
pass
@@ -0,0 +1,6 @@
from mod import Base, default
class Sub(Base):
def method(self, param=default):
super().method(param)
@@ -0,0 +1,6 @@
default = object()
class Base:
def method(self, param=default):
pass
@@ -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";