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:
Marcus Mews
2025-08-22 14:18:05 +00:00
committed by intellij-monorepo-bot
parent 22de1f427f
commit 992a8cdfb5
57 changed files with 492 additions and 89 deletions
@@ -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;
}
@@ -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) {
@@ -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);
}
}
@@ -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() {
@@ -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;
}
}
@@ -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;
}
}
@@ -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;
}
}
@@ -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,4 +1,4 @@
from typing import Dict, Any
from typing import Any, Dict
def func(x):
@@ -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]: ...
@@ -0,0 +1,3 @@
from m import foo_bar
v<caret>ar = foo_bar()
@@ -0,0 +1,5 @@
from typing import Literal
from m import foo_bar, MyEnum
var: [Literal[MyEnum.A]] = foo_bar()
@@ -1,4 +1,4 @@
from typing import List, Any
from typing import Any, List
def func():
@@ -1,4 +1,4 @@
from typing import Coroutine, Any
from typing import Any, Coroutine
async def bar() -> int:
@@ -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)
@@ -1,3 +1,5 @@
from typing import Any
x = 42
# print('commented')
@@ -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):
@@ -1,3 +1,6 @@
from typing import Any
class Test:
def method(self, x):
def func():
@@ -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))
@@ -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 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
@@ -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");