mirror of
https://gitflic.ru/project/openide/openide.git
synced 2026-09-27 10:03:11 +07:00
PY-83066 extract method doesn't add imports
- add imports for types of arguments and return values - sort new imports by name GitOrigin-RevId: 861c4733b47a4b056582b359c62c9f1360814178
This commit is contained in:
committed by
intellij-monorepo-bot
parent
22de1f427f
commit
992a8cdfb5
+33
-14
@@ -16,51 +16,70 @@
|
||||
package com.jetbrains.python.codeInsight.codeFragment;
|
||||
|
||||
import com.intellij.codeInsight.codeFragment.CodeFragment;
|
||||
import org.jetbrains.annotations.NotNull;
|
||||
import com.intellij.openapi.util.Pair;
|
||||
import com.jetbrains.python.psi.types.PyType;
|
||||
import org.jetbrains.annotations.NotNullByDefault;
|
||||
import org.jetbrains.annotations.Nullable;
|
||||
|
||||
import java.util.Map;
|
||||
import java.util.Set;
|
||||
|
||||
@NotNullByDefault
|
||||
public class PyCodeFragment extends CodeFragment {
|
||||
private final @NotNull Map<@NotNull String, @NotNull String> myInputTypes;
|
||||
/** Maps variable names to their type names and types. */
|
||||
private final Map<String, Pair<String, PyType>> myInputTypes;
|
||||
private final @Nullable String myOutputType;
|
||||
private final @NotNull Set<@NotNull String> myGlobalWrites;
|
||||
private final @NotNull Set<@NotNull String> myNonlocalWrites;
|
||||
private final Set<PyType> myOutputTypes;
|
||||
private final Set<String> myGlobalWrites;
|
||||
private final Set<String> myNonlocalWrites;
|
||||
private final boolean myYieldInside;
|
||||
private final boolean myAsync;
|
||||
|
||||
public PyCodeFragment(final @NotNull Set<@NotNull String> input,
|
||||
final @NotNull Set<@NotNull String> output,
|
||||
final @NotNull Map<@NotNull String, @NotNull String> inputTypes,
|
||||
public PyCodeFragment(final Set<String> input,
|
||||
final Set<String> output,
|
||||
final Map<String, Pair<String, PyType>> inputTypeNames,
|
||||
final @Nullable String outputType,
|
||||
final @NotNull Set<@NotNull String> globalWrites,
|
||||
final @NotNull Set<@NotNull String> nonlocalWrites,
|
||||
final Set<PyType> outputTypes,
|
||||
final Set<String> globalWrites,
|
||||
final Set<String> nonlocalWrites,
|
||||
final boolean returnInside,
|
||||
final boolean yieldInside,
|
||||
final boolean isAsync) {
|
||||
super(input, output, returnInside);
|
||||
myInputTypes = inputTypes;
|
||||
myInputTypes = inputTypeNames;
|
||||
myOutputType = outputType;
|
||||
myOutputTypes = outputTypes;
|
||||
myGlobalWrites = globalWrites;
|
||||
myNonlocalWrites = nonlocalWrites;
|
||||
myYieldInside = yieldInside;
|
||||
myAsync = isAsync;
|
||||
}
|
||||
|
||||
public @NotNull Map<@NotNull String, @NotNull String> getInputTypes() {
|
||||
return myInputTypes;
|
||||
/** Returns the type name of the input variable with the given name. */
|
||||
public @Nullable String getInputTypeName(String varName) {
|
||||
Pair<String, PyType> type = myInputTypes.get(varName);
|
||||
return type == null ? null : type.first;
|
||||
}
|
||||
|
||||
/** Returns the type of the input variable with the given name. */
|
||||
public @Nullable PyType getInputType(String varName) {
|
||||
Pair<String, PyType> type = myInputTypes.get(varName);
|
||||
return type == null ? null : type.second;
|
||||
}
|
||||
|
||||
public @Nullable String getOutputType() {
|
||||
return myOutputType;
|
||||
}
|
||||
|
||||
public @NotNull Set<@NotNull String> getGlobalWrites() {
|
||||
public Set<PyType> getOutputTypes() {
|
||||
return myOutputTypes;
|
||||
}
|
||||
|
||||
public Set<String> getGlobalWrites() {
|
||||
return myGlobalWrites;
|
||||
}
|
||||
|
||||
public @NotNull Set<@NotNull String> getNonlocalWrites() {
|
||||
public Set<String> getNonlocalWrites() {
|
||||
return myNonlocalWrites;
|
||||
}
|
||||
|
||||
|
||||
+23
-24
@@ -56,20 +56,27 @@ public final class PyCodeFragmentUtil {
|
||||
final Set<String> nonlocalWrites = getNonlocalWrites(subGraph, owner);
|
||||
|
||||
final TypeEvalContext context = TypeEvalContext.userInitiated(startInScope.getProject(), startInScope.getContainingFile());
|
||||
final Set<String> inputNames = new HashSet<>();
|
||||
final Map<String, String> inputTypeNames = new HashMap<>();
|
||||
final Set<String> inputNames = new LinkedHashSet<>();
|
||||
final Map<String, Pair<String, PyType>> inputTypes = new HashMap<>();
|
||||
for (PsiElement element : filterElementsInScope(getInputElements(subGraph, graph), owner)) {
|
||||
// Ignore "self" and "cls", they are generated automatically when extracting any method fragment
|
||||
if (resolvesToBoundMethodParameter(element)) {
|
||||
continue;
|
||||
}
|
||||
addNameReturnType(globalWrites, nonlocalWrites, element, inputNames, inputTypeNames, null, context);
|
||||
Pair<String, PyType> variable = getVariable(globalWrites, nonlocalWrites, element, inputNames, context);
|
||||
if (variable != null && variable.second != null) {
|
||||
String typeName = PythonDocumentationProvider.getTypeHint(variable.second, context);
|
||||
inputTypes.put(variable.first, Pair.create(typeName, variable.second));
|
||||
}
|
||||
}
|
||||
|
||||
final Set<String> outputNames = new HashSet<>();
|
||||
final Set<String> outputNames = new LinkedHashSet<>();
|
||||
final List<PyType> outputTypes = new ArrayList<>();
|
||||
for (PsiElement element : getOutputElements(subGraph, graph)) {
|
||||
addNameReturnType(globalWrites, nonlocalWrites, element, outputNames, null, outputTypes, context);
|
||||
Pair<String, PyType> variable = getVariable(globalWrites, nonlocalWrites, element, outputNames, context);
|
||||
if (variable != null) {
|
||||
outputTypes.add(variable.second);
|
||||
}
|
||||
}
|
||||
if (singleExpression != null) {
|
||||
PyType returnType = getType(singleExpression, context);
|
||||
@@ -88,30 +95,22 @@ public final class PyCodeFragmentUtil {
|
||||
}
|
||||
final boolean isAsync = owner instanceof PyFunction && ((PyFunction)owner).isAsync();
|
||||
|
||||
return new PyCodeFragment(inputNames, outputNames, inputTypeNames, outputTypeName, globalWrites, nonlocalWrites,
|
||||
subGraphAnalysis.returns > 0, yieldsFound, isAsync);
|
||||
return new PyCodeFragment(inputNames, outputNames, inputTypes, outputTypeName, new LinkedHashSet<>(outputTypes),
|
||||
globalWrites, nonlocalWrites, subGraphAnalysis.returns > 0, yieldsFound, isAsync);
|
||||
}
|
||||
|
||||
private static void addNameReturnType(@NotNull Set<String> globalWrites,
|
||||
@NotNull Set<String> nonlocalWrites,
|
||||
@NotNull PsiElement element,
|
||||
@NotNull Set<String> varNames,
|
||||
@Nullable Map<String, String> varTypeNames,
|
||||
@Nullable List<PyType> outputTypes,
|
||||
@NotNull TypeEvalContext context) {
|
||||
private static @Nullable Pair<String, PyType> getVariable(@NotNull Set<String> globalWrites,
|
||||
@NotNull Set<String> nonlocalWrites,
|
||||
@NotNull PsiElement element,
|
||||
@NotNull Set<String> variableNames,
|
||||
@NotNull TypeEvalContext context) {
|
||||
String name = getName(element);
|
||||
if (name == null || globalWrites.contains(name) || nonlocalWrites.contains(name) || varNames.contains(name)) {
|
||||
return;
|
||||
if (name == null || globalWrites.contains(name) || nonlocalWrites.contains(name) || variableNames.contains(name)) {
|
||||
return null;
|
||||
}
|
||||
varNames.add(name);
|
||||
PyType type = getType(element, context);
|
||||
if (varTypeNames != null) {
|
||||
String typeName = type == null ? null : PythonDocumentationProvider.getTypeHint(type, context);
|
||||
varTypeNames.put(name, typeName);
|
||||
}
|
||||
if (outputTypes != null) {
|
||||
outputTypes.add(type);
|
||||
}
|
||||
variableNames.add(name);
|
||||
return Pair.create(name, type);
|
||||
}
|
||||
|
||||
private static @Nullable PyType getType(@NotNull PsiElement element, @NotNull TypeEvalContext context) {
|
||||
|
||||
+5
-3
@@ -272,8 +272,10 @@ public final class PyTypeHintGenerationUtil {
|
||||
}
|
||||
}
|
||||
|
||||
public static void addImportsForTypeAnnotations(@NotNull List<String> types, @NotNull PsiElement anchor) {
|
||||
final Set<PsiNamedElement> symbols = new LinkedHashSet<>();
|
||||
/** Adds imports for type annotations. Sorts imports by name. */
|
||||
public static void addImportsForTypeAnnotations(@NotNull Collection<String> types, @NotNull PsiElement anchor) {
|
||||
final Set<PsiNamedElement> symbols =
|
||||
new TreeSet<>(Comparator.comparing(PsiNamedElement::getName, Comparator.nullsFirst(Comparator.naturalOrder())));
|
||||
|
||||
for (String type : types) {
|
||||
collectImportTargetsFromTypeExpression(type, anchor, symbols);
|
||||
@@ -287,7 +289,7 @@ public final class PyTypeHintGenerationUtil {
|
||||
|
||||
private static void collectImportTargetsFromTypeExpression(@NotNull String typeExpressionText,
|
||||
@NotNull PsiElement anchor,
|
||||
@NotNull Set<PsiNamedElement> symbols) {
|
||||
@NotNull Set<@NotNull PsiNamedElement> symbols) {
|
||||
PyExpression typeExpression = PyUtil.createExpressionFromFragment(typeExpressionText, anchor);
|
||||
assert typeExpression != null;
|
||||
PyQualifiedNameResolveContext qNameResolveContext = PyResolveImportUtil.fromFoothold(anchor);
|
||||
|
||||
@@ -18,6 +18,7 @@ import com.jetbrains.python.highlighting.PyHighlighter;
|
||||
import com.jetbrains.python.psi.LanguageLevel;
|
||||
import com.jetbrains.python.psi.PyExpression;
|
||||
import com.jetbrains.python.psi.PyQualifiedNameOwner;
|
||||
import com.jetbrains.python.psi.PyReferenceExpression;
|
||||
import com.jetbrains.python.psi.types.*;
|
||||
import one.util.streamex.StreamEx;
|
||||
import org.jetbrains.annotations.Nls;
|
||||
@@ -528,6 +529,27 @@ public abstract class PyTypeRenderer extends PyTypeVisitorExt<@NotNull HtmlChunk
|
||||
return result.toFragment();
|
||||
}
|
||||
|
||||
@Override
|
||||
public @NotNull HtmlChunk visitPyLiteralType(@NotNull PyLiteralType literalType) {
|
||||
HtmlBuilder result = new HtmlBuilder();
|
||||
result.append(HtmlChunk.raw(isRenderingFqn() ? "typing.Literal" : "Literal")); //NON-NLS
|
||||
result.append("[");
|
||||
@Nullable String classQName = literalType.getClassQName();
|
||||
if (isRenderingFqn() && classQName != null && literalType.getExpression() instanceof PyReferenceExpression refExpr) {
|
||||
result.append(classQName);
|
||||
if (refExpr.getName() != null) {
|
||||
result.append(".");
|
||||
result.append(refExpr.getName());
|
||||
}
|
||||
}
|
||||
else {
|
||||
String enumOrLiteral = StringUtil.notNullize(literalType.getExpression().getText()).trim();
|
||||
result.appendRaw(enumOrLiteral); // append raw since the literal can include quotes: Literal["foo"]
|
||||
}
|
||||
result.append("]");
|
||||
return result.toFragment();
|
||||
}
|
||||
|
||||
protected final @Nullable @NlsSafe String getTypeName(@NotNull PyType type) {
|
||||
if (isNoneType(type)) {
|
||||
return PyNames.NONE;
|
||||
|
||||
@@ -24,10 +24,7 @@ import com.jetbrains.python.psi.PyClass;
|
||||
import com.jetbrains.python.psi.PyPsiFacade;
|
||||
import com.jetbrains.python.psi.impl.PyBuiltinCache;
|
||||
import one.util.streamex.StreamEx;
|
||||
import org.jetbrains.annotations.ApiStatus;
|
||||
import org.jetbrains.annotations.Contract;
|
||||
import org.jetbrains.annotations.NotNull;
|
||||
import org.jetbrains.annotations.Nullable;
|
||||
import org.jetbrains.annotations.*;
|
||||
|
||||
import java.util.*;
|
||||
import java.util.function.BinaryOperator;
|
||||
@@ -215,4 +212,39 @@ public final class PyTypeUtil {
|
||||
public static boolean inheritsAny(@NotNull PyType type, @NotNull TypeEvalContext context) {
|
||||
return type instanceof PyClassLikeType classLikeType && classLikeType.getAncestorTypes(context).contains(null);
|
||||
}
|
||||
|
||||
/**
|
||||
* Collects a set of types that participate in the textual type hint representation of {@code type}.
|
||||
* The returned set preserves a stable DFS order and is unmodifiable.
|
||||
*/
|
||||
public static @NotNull @UnmodifiableView Set<PyType> collectTypeComponentsFromType(@Nullable PyType type,
|
||||
@NotNull TypeEvalContext context) {
|
||||
Set<PyType> result = new LinkedHashSet<>();
|
||||
|
||||
PyRecursiveTypeVisitor.traverse(type, context, new PyRecursiveTypeVisitor.PyTypeTraverser() {
|
||||
@Override
|
||||
public @NotNull PyRecursiveTypeVisitor.Traversal visitPyType(@NotNull PyType pyType) {
|
||||
result.add(pyType);
|
||||
return super.visitPyType(pyType);
|
||||
}
|
||||
|
||||
@Override
|
||||
public PyRecursiveTypeVisitor.@NotNull Traversal visitPyLiteralType(@NotNull PyLiteralType literalType) {
|
||||
PyClassLikeType literalClassType = literalType.getPyClass().getType(context);
|
||||
if (literalClassType != null) {
|
||||
// Adds eg. signal.Handler when the given type was Literal[Handlers.SIG_DFL]
|
||||
result.add(literalClassType);
|
||||
}
|
||||
return super.visitPyLiteralType(literalType);
|
||||
}
|
||||
|
||||
@Override
|
||||
public PyRecursiveTypeVisitor.@NotNull Traversal visitUnknownType() {
|
||||
result.add(null); // add Any type
|
||||
return super.visitUnknownType();
|
||||
}
|
||||
});
|
||||
|
||||
return Collections.unmodifiableSet(result);
|
||||
}
|
||||
}
|
||||
|
||||
+6
-5
@@ -18,13 +18,14 @@ import com.jetbrains.python.refactoring.extractmethod.PyVariableData;
|
||||
import com.jetbrains.python.refactoring.introduce.IntroduceOperation;
|
||||
import com.jetbrains.python.refactoring.introduce.IntroduceValidator;
|
||||
import org.jetbrains.annotations.ApiStatus;
|
||||
import org.jetbrains.annotations.NotNull;
|
||||
import org.jetbrains.annotations.NotNullByDefault;
|
||||
import org.jetbrains.annotations.Nullable;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.function.Consumer;
|
||||
|
||||
@ApiStatus.Experimental
|
||||
@NotNullByDefault
|
||||
public class PyRefactoringUiService {
|
||||
public void performIntroduceWithDialog(IntroduceOperation operation,
|
||||
@NlsContexts.DialogTitle String dialogTitle,
|
||||
@@ -55,12 +56,12 @@ public class PyRefactoringUiService {
|
||||
final ExtractMethodDecorator<Object> decorator,
|
||||
final FileType type, String helpId) {
|
||||
return new PyExtractMethodSettings(defaultName, new PyVariableData[0], fragment.getOutputType(),
|
||||
PyExtractMethodUtil.getAddTypeAnnotations(project));
|
||||
fragment.getOutputTypes(), PyExtractMethodUtil.getAddTypeAnnotations(project));
|
||||
}
|
||||
|
||||
public void showPyInlineFunctionDialog(@NotNull Project project,
|
||||
@NotNull Editor editor,
|
||||
@NotNull PyFunction function, @Nullable PsiReference reference) {
|
||||
public void showPyInlineFunctionDialog(Project project,
|
||||
Editor editor,
|
||||
PyFunction function, @Nullable PsiReference reference) {
|
||||
}
|
||||
|
||||
public static PyRefactoringUiService getInstance() {
|
||||
|
||||
+33
-9
@@ -1,39 +1,52 @@
|
||||
package com.jetbrains.python.refactoring.extractmethod;
|
||||
|
||||
import com.intellij.refactoring.extractMethod.ExtractMethodSettings;
|
||||
import org.jetbrains.annotations.NotNull;
|
||||
import com.jetbrains.python.psi.types.PyType;
|
||||
import org.jetbrains.annotations.NotNullByDefault;
|
||||
import org.jetbrains.annotations.Nullable;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
import java.util.Set;
|
||||
|
||||
@NotNullByDefault
|
||||
public class PyExtractMethodSettings implements ExtractMethodSettings<Object> {
|
||||
private final String myMethodName;
|
||||
private final PyVariableData @NotNull [] myVariableData;
|
||||
private final String myReturnTypeName;
|
||||
private final PyVariableData[] myVariableData;
|
||||
private final @Nullable String myReturnTypeName;
|
||||
private final Set<PyType> myReturnTypes;
|
||||
private final boolean myUseTypeAnnotations;
|
||||
|
||||
public PyExtractMethodSettings(@NotNull String methodName,
|
||||
PyVariableData @NotNull [] variableData,
|
||||
String returnTypeName,
|
||||
public PyExtractMethodSettings(String methodName,
|
||||
PyVariableData[] variableData,
|
||||
@Nullable String returnTypeName,
|
||||
Set<PyType> returnTypes,
|
||||
boolean useTypeAnnotations) {
|
||||
myMethodName = methodName;
|
||||
myVariableData = variableData;
|
||||
myReturnTypeName = returnTypeName;
|
||||
myReturnTypes = returnTypes;
|
||||
myUseTypeAnnotations = useTypeAnnotations;
|
||||
}
|
||||
|
||||
@Override
|
||||
public @NotNull String getMethodName() {
|
||||
public String getMethodName() {
|
||||
return myMethodName;
|
||||
}
|
||||
|
||||
@Override
|
||||
public PyVariableData @NotNull [] getAbstractVariableData() {
|
||||
public PyVariableData[] getAbstractVariableData() {
|
||||
return myVariableData;
|
||||
}
|
||||
|
||||
public String getReturnTypeName() {
|
||||
public @Nullable String getReturnTypeName() {
|
||||
return myReturnTypeName;
|
||||
}
|
||||
|
||||
public Set<PyType> getReturnTypeFqns() {
|
||||
return myReturnTypes;
|
||||
}
|
||||
|
||||
public boolean isUseTypeAnnotations() {
|
||||
return myUseTypeAnnotations;
|
||||
}
|
||||
@@ -42,4 +55,15 @@ public class PyExtractMethodSettings implements ExtractMethodSettings<Object> {
|
||||
public @Nullable Object getVisibility() {
|
||||
return null;
|
||||
}
|
||||
|
||||
List<PyType> getAllTypes() {
|
||||
List<PyType> result = new ArrayList<>();
|
||||
for (PyVariableData variableData : myVariableData) {
|
||||
if (variableData.type != null) {
|
||||
result.add(variableData.type);
|
||||
}
|
||||
}
|
||||
result.addAll(myReturnTypes);
|
||||
return result;
|
||||
}
|
||||
}
|
||||
|
||||
+55
-18
@@ -13,6 +13,7 @@ import com.intellij.openapi.ui.MessageDialogBuilder;
|
||||
import com.intellij.openapi.util.Pair;
|
||||
import com.intellij.openapi.util.TextRange;
|
||||
import com.intellij.openapi.util.text.StringUtil;
|
||||
import com.intellij.openapi.util.text.Strings;
|
||||
import com.intellij.psi.*;
|
||||
import com.intellij.psi.impl.source.codeStyle.CodeEditUtil;
|
||||
import com.intellij.psi.util.PsiTreeUtil;
|
||||
@@ -38,9 +39,13 @@ import com.jetbrains.python.codeInsight.controlflow.ControlFlowCache;
|
||||
import com.jetbrains.python.codeInsight.controlflow.ScopeOwner;
|
||||
import com.jetbrains.python.codeInsight.dataflow.scope.Scope;
|
||||
import com.jetbrains.python.codeInsight.dataflow.scope.ScopeUtil;
|
||||
import com.jetbrains.python.codeInsight.intentions.PyTypeHintGenerationUtil;
|
||||
import com.jetbrains.python.documentation.PythonDocumentationProvider;
|
||||
import com.jetbrains.python.psi.*;
|
||||
import com.jetbrains.python.psi.impl.PyFunctionBuilder;
|
||||
import com.jetbrains.python.psi.impl.PyPsiUtils;
|
||||
import com.jetbrains.python.psi.types.PyType;
|
||||
import com.jetbrains.python.psi.types.TypeEvalContext;
|
||||
import com.jetbrains.python.refactoring.PyRefactoringUiService;
|
||||
import com.jetbrains.python.refactoring.PyReplaceExpressionUtil;
|
||||
import org.jetbrains.annotations.NotNull;
|
||||
@@ -48,6 +53,8 @@ import org.jetbrains.annotations.Nullable;
|
||||
|
||||
import java.util.*;
|
||||
|
||||
import static com.jetbrains.python.psi.types.PyTypeUtil.collectTypeComponentsFromType;
|
||||
|
||||
public final class PyExtractMethodUtil {
|
||||
public static final String NAME = "extract.method.name";
|
||||
private static final String ADD_TYPE_ANNOTATIONS_VALUE_KEY = "settings.extract.method.addTypeAnnotations";
|
||||
@@ -170,21 +177,16 @@ public final class PyExtractMethodUtil {
|
||||
final List<SimpleMatch> duplicates = collectDuplicates(finder, statement1, insertedMethod);
|
||||
|
||||
// replace statements with call
|
||||
PsiElement insertedCallElement = WriteAction.compute(() -> replaceElements(elementsRange, callElement));
|
||||
insertedCallElement = CodeInsightUtilCore.forcePsiPostprocessAndRestoreElement(insertedCallElement);
|
||||
PsiElement insertedCallElement = CodeInsightUtilCore.forcePsiPostprocessAndRestoreElement(WriteAction.compute(
|
||||
() -> replaceElements(elementsRange, callElement)));
|
||||
|
||||
SmartPointerManager pointerManager = SmartPointerManager.getInstance(project);
|
||||
if (processDuplicates) {
|
||||
pointers.addAll(ContainerUtil.map(duplicates, p -> pointerManager.createSmartPsiFileRangePointer(file, p.getStartElement().getTextRange())));
|
||||
}
|
||||
|
||||
if (insertedCallElement != null) {
|
||||
pointers.add(0, pointerManager.createSmartPsiFileRangePointer(file, insertedMethod.getNameIdentifier().getTextRange()));
|
||||
pointers.add(pointerManager.createSmartPsiFileRangePointer(file, insertedCallElement.getTextRange()));
|
||||
if (processDuplicates) {
|
||||
processDuplicates(duplicates, insertedCallElement, editor);
|
||||
}
|
||||
}
|
||||
processDuplicatesAndAddImports(project, editor, processDuplicates, pointers, methodSettings, insertedMethod,
|
||||
duplicates, insertedCallElement, file, pointerManager);
|
||||
|
||||
// Set editor
|
||||
setSelectionAndCaret(editor, insertedCallElement);
|
||||
@@ -381,13 +383,8 @@ public final class PyExtractMethodUtil {
|
||||
}
|
||||
if (callElement != null) {
|
||||
insertedCallElement = WriteAction.compute(() -> PyReplaceExpressionUtil.replaceExpression(expression, callElement));
|
||||
if (insertedCallElement != null) {
|
||||
pointers.add(0, pointerManager.createSmartPsiFileRangePointer(file, insertedMethod.getNameIdentifier().getTextRange()));
|
||||
pointers.add(pointerManager.createSmartPsiFileRangePointer(file, insertedCallElement.getTextRange()));
|
||||
if (processDuplicates) {
|
||||
processDuplicates(duplicates, insertedCallElement, editor);
|
||||
}
|
||||
}
|
||||
processDuplicatesAndAddImports(project, editor, processDuplicates, pointers, methodSettings, insertedMethod,
|
||||
duplicates, insertedCallElement, file, pointerManager);
|
||||
}
|
||||
setSelectionAndCaret(editor, insertedCallElement);
|
||||
// Set editor
|
||||
@@ -396,6 +393,44 @@ public final class PyExtractMethodUtil {
|
||||
return pointers;
|
||||
}
|
||||
|
||||
private static void processDuplicatesAndAddImports(@NotNull Project project,
|
||||
@NotNull Editor editor,
|
||||
@NotNull Boolean processDuplicates,
|
||||
@NotNull List<SmartPsiFileRange> pointers,
|
||||
@NotNull PyExtractMethodSettings methodSettings,
|
||||
@NotNull PyFunction insertedMethod,
|
||||
@NotNull List<SimpleMatch> duplicates,
|
||||
PsiElement insertedCallElement,
|
||||
@NotNull PsiFile file,
|
||||
@NotNull SmartPointerManager pointerManager) {
|
||||
|
||||
if (insertedCallElement == null) {
|
||||
return;
|
||||
}
|
||||
pointers.add(0, pointerManager.createSmartPsiFileRangePointer(file, insertedMethod.getNameIdentifier().getTextRange()));
|
||||
pointers.add(pointerManager.createSmartPsiFileRangePointer(file, insertedCallElement.getTextRange()));
|
||||
if (processDuplicates) {
|
||||
processDuplicates(duplicates, insertedCallElement, editor);
|
||||
}
|
||||
if (getAddTypeAnnotations(project)) {
|
||||
TypeEvalContext context = TypeEvalContext.userInitiated(project, file);
|
||||
Set<String> allTypesAsStrings = new HashSet<>();
|
||||
for (PyType type : methodSettings.getAllTypes()) {
|
||||
for (PyType type2 : collectTypeComponentsFromType(type, context)) {
|
||||
if (type2 == null || type2.getDeclarationElement() == null || type2.getDeclarationElement().isValid()) {
|
||||
String typeFqn = PythonDocumentationProvider.getFullyQualifiedTypeHint(type2, context);
|
||||
if (Strings.isNotEmpty(typeFqn)) {
|
||||
allTypesAsStrings.add(typeFqn);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
WriteAction.run(() -> {
|
||||
PyTypeHintGenerationUtil.addImportsForTypeAnnotations(allTypesAsStrings, insertedMethod);
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
private static void setSelectionAndCaret(@NotNull Editor editor, final @Nullable PsiElement callElement) {
|
||||
editor.getSelectionModel().removeSelection();
|
||||
if (callElement != null) {
|
||||
@@ -631,11 +666,12 @@ public final class PyExtractMethodUtil {
|
||||
d.name = in + "_new";
|
||||
d.originalName = in;
|
||||
d.passAsParameter = true;
|
||||
d.typeName = fragment.getInputTypes().get(in);
|
||||
d.typeName = fragment.getInputTypeName(in);
|
||||
d.type = fragment.getInputType(in);
|
||||
data.add(d);
|
||||
}
|
||||
return new PyExtractMethodSettings(name, data.toArray(new PyVariableData[0]), fragment.getOutputType(),
|
||||
getAddTypeAnnotations(project));
|
||||
fragment.getOutputTypes(), getAddTypeAnnotations(project));
|
||||
}
|
||||
|
||||
final boolean isMethod = PyPsiUtils.isMethodContext(element);
|
||||
@@ -727,4 +763,5 @@ public final class PyExtractMethodUtil {
|
||||
boolean selected = PropertiesComponent.getInstance(project).getBoolean(ADD_TYPE_ANNOTATIONS_VALUE_KEY, ADD_TYPE_ANNOTATIONS_DEFAULT);
|
||||
return selected;
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
+6
@@ -1,13 +1,19 @@
|
||||
package com.jetbrains.python.refactoring.extractmethod;
|
||||
|
||||
import com.intellij.refactoring.util.AbstractVariableData;
|
||||
import com.jetbrains.python.psi.types.PyType;
|
||||
import org.jetbrains.annotations.Nullable;
|
||||
|
||||
public class PyVariableData extends AbstractVariableData {
|
||||
public @Nullable String typeName;
|
||||
public @Nullable PyType type;
|
||||
|
||||
|
||||
public @Nullable String getTypeName() {
|
||||
return typeName;
|
||||
}
|
||||
|
||||
public @Nullable PyType getType() {
|
||||
return type;
|
||||
}
|
||||
}
|
||||
|
||||
-1
@@ -43,7 +43,6 @@ import com.jetbrains.python.PyPsiBundle;
|
||||
import com.jetbrains.python.PyTokenTypes;
|
||||
import com.jetbrains.python.codeInsight.dataflow.scope.ScopeUtil;
|
||||
import com.jetbrains.python.psi.*;
|
||||
import com.jetbrains.python.psi.impl.PyBuiltinCache;
|
||||
import com.jetbrains.python.psi.impl.PyPsiUtils;
|
||||
import com.jetbrains.python.psi.resolve.PyResolveContext;
|
||||
import com.jetbrains.python.psi.resolve.PyResolveUtil;
|
||||
|
||||
@@ -21,6 +21,7 @@ import org.jetbrains.annotations.Nullable;
|
||||
import javax.swing.*;
|
||||
import java.awt.*;
|
||||
import java.util.List;
|
||||
import java.util.Objects;
|
||||
import java.util.function.Predicate;
|
||||
|
||||
public class PyExtractMethodDialog extends AbstractExtractMethodDialog<Object> {
|
||||
@@ -62,7 +63,7 @@ public class PyExtractMethodDialog extends AbstractExtractMethodDialog<Object> {
|
||||
@NotNull
|
||||
public PyExtractMethodSettings getExtractMethodSettings() {
|
||||
return new PyExtractMethodSettings(getMethodName(), getAbstractVariableData(), ((PyCodeFragment)myFragment).getOutputType(),
|
||||
myAddTypeAnnotationsCheckbox.isSelected());
|
||||
((PyCodeFragment)myFragment).getOutputTypes(), myAddTypeAnnotationsCheckbox.isSelected());
|
||||
}
|
||||
|
||||
@Override
|
||||
@@ -75,7 +76,8 @@ public class PyExtractMethodDialog extends AbstractExtractMethodDialog<Object> {
|
||||
data.originalName = name;
|
||||
data.name = name;
|
||||
data.passAsParameter = true;
|
||||
data.typeName = ((PyCodeFragment)myFragment).getInputTypes().get(name);
|
||||
data.typeName = ((PyCodeFragment)myFragment).getInputTypeName(name);
|
||||
data.type = ((PyCodeFragment)myFragment).getInputType(name);
|
||||
datas[i] = data;
|
||||
}
|
||||
return datas;
|
||||
@@ -110,8 +112,9 @@ public class PyExtractMethodDialog extends AbstractExtractMethodDialog<Object> {
|
||||
|
||||
@Override
|
||||
public void setValue(@NotNull PyVariableData data, @NotNull String value) {
|
||||
if (myNameValidator.test(value)) {
|
||||
if (myNameValidator.test(value) && !Objects.equals(data.getTypeName(), value)) {
|
||||
data.typeName = value;
|
||||
data.type = null; // the user needs to import the type he specified manually
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -20,12 +20,13 @@ import com.jetbrains.python.refactoring.inline.PyInlineFunctionDialog;
|
||||
import com.jetbrains.python.refactoring.introduce.IntroduceOperation;
|
||||
import com.jetbrains.python.refactoring.introduce.IntroduceValidator;
|
||||
import com.jetbrains.python.refactoring.introduce.PyIntroduceHandlerUi;
|
||||
import org.jetbrains.annotations.NotNull;
|
||||
import org.jetbrains.annotations.NotNullByDefault;
|
||||
import org.jetbrains.annotations.Nullable;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.function.Consumer;
|
||||
|
||||
@NotNullByDefault
|
||||
public final class PyRefactoringUiServiceImpl extends PyRefactoringUiService {
|
||||
@Override
|
||||
public void showIntroduceTargetChooser(IntroduceOperation operation,
|
||||
@@ -92,9 +93,9 @@ public final class PyRefactoringUiServiceImpl extends PyRefactoringUiService {
|
||||
}
|
||||
|
||||
@Override
|
||||
public void showPyInlineFunctionDialog(@NotNull Project project,
|
||||
@NotNull Editor editor,
|
||||
@NotNull PyFunction function,
|
||||
public void showPyInlineFunctionDialog(Project project,
|
||||
Editor editor,
|
||||
PyFunction function,
|
||||
@Nullable PsiReference reference) {
|
||||
new PyInlineFunctionDialog(project, editor, function, reference).show();
|
||||
}
|
||||
|
||||
+1
-1
@@ -1,4 +1,4 @@
|
||||
from typing import Dict, Any
|
||||
from typing import Any, Dict
|
||||
|
||||
|
||||
def func(x):
|
||||
|
||||
+10
@@ -0,0 +1,10 @@
|
||||
from enum import EnumType
|
||||
from typing import Literal
|
||||
|
||||
|
||||
class MyEnum(metaclass=EnumType):
|
||||
A = 1
|
||||
B = 2
|
||||
|
||||
|
||||
def foo_bar() -> Literal[MyEnum.A]: ...
|
||||
+3
@@ -0,0 +1,3 @@
|
||||
from m import foo_bar
|
||||
|
||||
v<caret>ar = foo_bar()
|
||||
+5
@@ -0,0 +1,5 @@
|
||||
from typing import Literal
|
||||
|
||||
from m import foo_bar, MyEnum
|
||||
|
||||
var: [Literal[MyEnum.A]] = foo_bar()
|
||||
+1
-1
@@ -1,4 +1,4 @@
|
||||
from typing import List, Any
|
||||
from typing import Any, List
|
||||
|
||||
|
||||
def func():
|
||||
|
||||
+1
-1
@@ -1,4 +1,4 @@
|
||||
from typing import Coroutine, Any
|
||||
from typing import Any, Coroutine
|
||||
|
||||
|
||||
async def bar() -> int:
|
||||
|
||||
+1
-1
@@ -1,4 +1,4 @@
|
||||
from typing import AsyncGenerator, Any
|
||||
from typing import Any, AsyncGenerator
|
||||
|
||||
|
||||
async def gen() -> AsyncGenerator[str | float, Any]:
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
from typing import Any
|
||||
|
||||
|
||||
async def foo(x):
|
||||
y = await bar(x)
|
||||
return await y
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
from typing import Any
|
||||
|
||||
|
||||
async def foo(x):
|
||||
y = await bar(x)
|
||||
return y
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
from typing import Any
|
||||
|
||||
|
||||
def foo() -> Any:
|
||||
return bbb
|
||||
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
from typing import Any
|
||||
|
||||
|
||||
def cylinder_volume(r, h):
|
||||
h * bar(r)
|
||||
|
||||
|
||||
+2
@@ -1,3 +1,5 @@
|
||||
from typing import Any
|
||||
|
||||
x = 42
|
||||
|
||||
# print('commented')
|
||||
|
||||
+3
@@ -1,3 +1,6 @@
|
||||
from typing import Any
|
||||
|
||||
|
||||
def f(n):
|
||||
return n * 2 if bar(n) else n + 1
|
||||
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
from typing import Any
|
||||
|
||||
|
||||
def f():
|
||||
a = do_smth()
|
||||
b1 = foo(a)
|
||||
|
||||
@@ -1,3 +1,7 @@
|
||||
from io import TextIOWrapper, _WrappedBuffer
|
||||
from typing import Any, IO
|
||||
|
||||
|
||||
def foo():
|
||||
for arg in sys.argv[1:]:
|
||||
try:
|
||||
|
||||
@@ -0,0 +1,8 @@
|
||||
def f(a):
|
||||
compiled = compile("x = 42", "<string>", "exec")
|
||||
body(compiled)
|
||||
|
||||
|
||||
def body(compiled_new):
|
||||
1
|
||||
compiled_new
|
||||
@@ -0,0 +1,11 @@
|
||||
from types import CodeType
|
||||
|
||||
|
||||
def f(a):
|
||||
compiled = compile("x = 42", "<string>", "exec")
|
||||
body(compiled)
|
||||
|
||||
|
||||
def body(compiled_new: CodeType):
|
||||
1
|
||||
compiled_new
|
||||
@@ -0,0 +1,4 @@
|
||||
def f(a):
|
||||
compiled = compile("x = 42", "<string>", "exec")
|
||||
<selection>1
|
||||
compiled</selection>
|
||||
@@ -0,0 +1,8 @@
|
||||
def f(a):
|
||||
file = open("test.txt", "w")
|
||||
body(file)
|
||||
|
||||
|
||||
def body(file_new):
|
||||
1
|
||||
file_new
|
||||
@@ -0,0 +1,11 @@
|
||||
from io import TextIOWrapper, _WrappedBuffer
|
||||
|
||||
|
||||
def f(a):
|
||||
file = open("test.txt", "w")
|
||||
body(file)
|
||||
|
||||
|
||||
def body(file_new: TextIOWrapper[_WrappedBuffer]):
|
||||
1
|
||||
file_new
|
||||
@@ -0,0 +1,4 @@
|
||||
def f(a):
|
||||
file = open("test.txt", "w")
|
||||
<selection>1
|
||||
file</selection>
|
||||
@@ -0,0 +1,8 @@
|
||||
def f(a):
|
||||
if a is 1:
|
||||
body(a)
|
||||
|
||||
|
||||
def body(a_new):
|
||||
1
|
||||
a_new
|
||||
@@ -0,0 +1,11 @@
|
||||
from typing import Literal
|
||||
|
||||
|
||||
def f(a):
|
||||
if a is 1:
|
||||
body(a)
|
||||
|
||||
|
||||
def body(a_new: Literal[1]):
|
||||
1
|
||||
a_new
|
||||
@@ -0,0 +1,4 @@
|
||||
def f(a):
|
||||
if a is 1:
|
||||
<selection>1
|
||||
a</selection>
|
||||
@@ -0,0 +1,15 @@
|
||||
from enum import Enum
|
||||
|
||||
class Color(Enum):
|
||||
RED = 1
|
||||
GREEN = 2
|
||||
BLUE = 3
|
||||
|
||||
def f(color):
|
||||
if color == Color.RED:
|
||||
body(color)
|
||||
|
||||
|
||||
def body(color_new):
|
||||
1
|
||||
color_new
|
||||
@@ -0,0 +1,17 @@
|
||||
from enum import Enum
|
||||
from typing import Literal
|
||||
|
||||
|
||||
class Color(Enum):
|
||||
RED = 1
|
||||
GREEN = 2
|
||||
BLUE = 3
|
||||
|
||||
def f(color):
|
||||
if color == Color.RED:
|
||||
body(color)
|
||||
|
||||
|
||||
def body(color_new: Literal[Color.RED]):
|
||||
1
|
||||
color_new
|
||||
@@ -0,0 +1,11 @@
|
||||
from enum import Enum
|
||||
|
||||
class Color(Enum):
|
||||
RED = 1
|
||||
GREEN = 2
|
||||
BLUE = 3
|
||||
|
||||
def f(color):
|
||||
if color == Color.RED:
|
||||
<selection>1
|
||||
color</selection>
|
||||
@@ -0,0 +1,11 @@
|
||||
import signal
|
||||
|
||||
|
||||
def f(sign) :
|
||||
if sign is signal.Handlers.SIG_DFL:
|
||||
body(sign)
|
||||
|
||||
|
||||
def body(sign_new):
|
||||
1
|
||||
sign_new
|
||||
@@ -0,0 +1,13 @@
|
||||
import signal
|
||||
from signal import Handlers
|
||||
from typing import Literal
|
||||
|
||||
|
||||
def f(sign) :
|
||||
if sign is signal.Handlers.SIG_DFL:
|
||||
body(sign)
|
||||
|
||||
|
||||
def body(sign_new: Literal[Handlers.SIG_DFL]):
|
||||
1
|
||||
sign_new
|
||||
@@ -0,0 +1,7 @@
|
||||
import signal
|
||||
|
||||
|
||||
def f(sign) :
|
||||
if sign is signal.Handlers.SIG_DFL:
|
||||
<selection>1
|
||||
sign</selection>
|
||||
@@ -1,3 +1,6 @@
|
||||
from typing import Any
|
||||
|
||||
|
||||
def foo(some_var):
|
||||
if bar(some_var):
|
||||
print('w00t')
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
from typing import Any
|
||||
|
||||
|
||||
def foo(some_var):
|
||||
if bar(some_var):
|
||||
print('w00t')
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
from typing import Any, Callable
|
||||
|
||||
|
||||
def foo():
|
||||
def f(x):
|
||||
return x
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
from typing import Any
|
||||
|
||||
|
||||
class Test:
|
||||
a = 5
|
||||
def method(self, b):
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
from typing import Any
|
||||
|
||||
|
||||
class Test:
|
||||
def method(self, a):
|
||||
def func(b):
|
||||
|
||||
+3
@@ -1,3 +1,6 @@
|
||||
from typing import Any
|
||||
|
||||
|
||||
class Test:
|
||||
def method(self, x):
|
||||
def func():
|
||||
|
||||
+3
@@ -1,3 +1,6 @@
|
||||
from typing import Any
|
||||
|
||||
|
||||
class Test:
|
||||
def method(self):
|
||||
def func(x):
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
from typing import Any
|
||||
|
||||
|
||||
def x(p_name, params):
|
||||
return bar(p_name, params), None
|
||||
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
from typing import Any
|
||||
|
||||
|
||||
def compound_duplicate(p1, p2):
|
||||
print(bar(p1))
|
||||
print(bar(p2))
|
||||
|
||||
+3
@@ -1,3 +1,6 @@
|
||||
from typing import Any
|
||||
|
||||
|
||||
def long_function_name(**kwargs): ...
|
||||
|
||||
def example_function():
|
||||
|
||||
+3
@@ -1,3 +1,6 @@
|
||||
from typing import Any
|
||||
|
||||
|
||||
def long_function_name(**kwargs): ...
|
||||
|
||||
def example_function():
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
from typing import Any
|
||||
|
||||
|
||||
def foo(f):
|
||||
x = 1
|
||||
x = bar(f, x)
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
from typing import Any
|
||||
|
||||
|
||||
def f(x, y):
|
||||
yield 'foo'
|
||||
return x, y
|
||||
|
||||
+9
@@ -295,6 +295,11 @@ public class PyAnnotateVariableTypeIntentionTest extends PyIntentionTestCase {
|
||||
doAnnotationTest();
|
||||
}
|
||||
|
||||
// PY-83066
|
||||
public void testAnnotationLiteralEnumType() {
|
||||
doMultiFileAnnotationTest(LanguageLevel.getLatest());
|
||||
}
|
||||
|
||||
// PY-46546
|
||||
public void testAnnotationGenericBuiltinList() {
|
||||
doTest(LanguageLevel.getLatest());
|
||||
@@ -342,6 +347,10 @@ public class PyAnnotateVariableTypeIntentionTest extends PyIntentionTestCase {
|
||||
runWithLanguageLevel(LanguageLevel.PYTHON36, () -> doMultiFileTest(PyPsiBundle.message("INTN.NAME.add.type.hint.for.variable")));
|
||||
}
|
||||
|
||||
public void doMultiFileAnnotationTest(LanguageLevel languageLevel) {
|
||||
runWithLanguageLevel(languageLevel, () -> doMultiFileTest(PyPsiBundle.message("INTN.NAME.add.type.hint.for.variable")));
|
||||
}
|
||||
|
||||
private void doMultiFileTest(@NotNull String hint) {
|
||||
myFixture.copyDirectoryToProject(getTestName(false), "");
|
||||
myFixture.configureByFile("main.py");
|
||||
|
||||
@@ -349,6 +349,31 @@ public class PyExtractMethodTest extends LightMarkedTestCase {
|
||||
doTest("body");
|
||||
}
|
||||
|
||||
// PY-83066
|
||||
public void testExtractAddsImport1() {
|
||||
doTest("body");
|
||||
}
|
||||
|
||||
// PY-83066
|
||||
public void testExtractAddsImport2() {
|
||||
doTest("body");
|
||||
}
|
||||
|
||||
// PY-83066
|
||||
public void testExtractAddsImport3() {
|
||||
doTest("body");
|
||||
}
|
||||
|
||||
// PY-83066
|
||||
public void testExtractAddsImport4() {
|
||||
doTest("body");
|
||||
}
|
||||
|
||||
// PY-83066
|
||||
public void testExtractAddsImport5() {
|
||||
doTest("body");
|
||||
}
|
||||
|
||||
// PY-35287
|
||||
public void testTypedStatements() {
|
||||
doTest("greeting");
|
||||
|
||||
Reference in New Issue
Block a user