mirror of
https://gitflic.ru/project/openide/openide.git
synced 2026-09-27 10:03:11 +07:00
PY-45588 Generalize PyOverrideImplementUtil.addImports to cover parameter defaults
GitOrigin-RevId: 828305e2a7449704c08503d8d68394d28ce98c83
This commit is contained in:
committed by
intellij-monorepo-bot
parent
b50c617b7c
commit
c7bba04743
@@ -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";
|
||||
|
||||
Reference in New Issue
Block a user