diff --git a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/codeFragment/PyCodeFragment.java b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/codeFragment/PyCodeFragment.java index 7ef15c24bb9c..d0a1807864d2 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/codeFragment/PyCodeFragment.java +++ b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/codeFragment/PyCodeFragment.java @@ -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> myInputTypes; private final @Nullable String myOutputType; - private final @NotNull Set<@NotNull String> myGlobalWrites; - private final @NotNull Set<@NotNull String> myNonlocalWrites; + private final Set myOutputTypes; + private final Set myGlobalWrites; + private final Set 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 input, + final Set output, + final Map> inputTypeNames, final @Nullable String outputType, - final @NotNull Set<@NotNull String> globalWrites, - final @NotNull Set<@NotNull String> nonlocalWrites, + final Set outputTypes, + final Set globalWrites, + final Set 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 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 type = myInputTypes.get(varName); + return type == null ? null : type.second; } public @Nullable String getOutputType() { return myOutputType; } - public @NotNull Set<@NotNull String> getGlobalWrites() { + public Set getOutputTypes() { + return myOutputTypes; + } + + public Set getGlobalWrites() { return myGlobalWrites; } - public @NotNull Set<@NotNull String> getNonlocalWrites() { + public Set getNonlocalWrites() { return myNonlocalWrites; } diff --git a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/codeFragment/PyCodeFragmentUtil.java b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/codeFragment/PyCodeFragmentUtil.java index 6bd7aa570b37..d6e92e8d85b1 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/codeFragment/PyCodeFragmentUtil.java +++ b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/codeFragment/PyCodeFragmentUtil.java @@ -56,20 +56,27 @@ public final class PyCodeFragmentUtil { final Set nonlocalWrites = getNonlocalWrites(subGraph, owner); final TypeEvalContext context = TypeEvalContext.userInitiated(startInScope.getProject(), startInScope.getContainingFile()); - final Set inputNames = new HashSet<>(); - final Map inputTypeNames = new HashMap<>(); + final Set inputNames = new LinkedHashSet<>(); + final Map> 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 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 outputNames = new HashSet<>(); + final Set outputNames = new LinkedHashSet<>(); final List outputTypes = new ArrayList<>(); for (PsiElement element : getOutputElements(subGraph, graph)) { - addNameReturnType(globalWrites, nonlocalWrites, element, outputNames, null, outputTypes, context); + Pair 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 globalWrites, - @NotNull Set nonlocalWrites, - @NotNull PsiElement element, - @NotNull Set varNames, - @Nullable Map varTypeNames, - @Nullable List outputTypes, - @NotNull TypeEvalContext context) { + private static @Nullable Pair getVariable(@NotNull Set globalWrites, + @NotNull Set nonlocalWrites, + @NotNull PsiElement element, + @NotNull Set 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) { diff --git a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/intentions/PyTypeHintGenerationUtil.java b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/intentions/PyTypeHintGenerationUtil.java index 9bb5e6270d27..bc478e6ddd2c 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/intentions/PyTypeHintGenerationUtil.java +++ b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/intentions/PyTypeHintGenerationUtil.java @@ -272,8 +272,10 @@ public final class PyTypeHintGenerationUtil { } } - public static void addImportsForTypeAnnotations(@NotNull List types, @NotNull PsiElement anchor) { - final Set symbols = new LinkedHashSet<>(); + /** Adds imports for type annotations. Sorts imports by name. */ + public static void addImportsForTypeAnnotations(@NotNull Collection types, @NotNull PsiElement anchor) { + final Set 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 symbols) { + @NotNull Set<@NotNull PsiNamedElement> symbols) { PyExpression typeExpression = PyUtil.createExpressionFromFragment(typeExpressionText, anchor); assert typeExpression != null; PyQualifiedNameResolveContext qNameResolveContext = PyResolveImportUtil.fromFoothold(anchor); diff --git a/python/python-psi-impl/src/com/jetbrains/python/documentation/PyTypeRenderer.java b/python/python-psi-impl/src/com/jetbrains/python/documentation/PyTypeRenderer.java index f42ff4b17ab4..5aad66cf210e 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/documentation/PyTypeRenderer.java +++ b/python/python-psi-impl/src/com/jetbrains/python/documentation/PyTypeRenderer.java @@ -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; diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeUtil.java b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeUtil.java index 628d564f6ffc..7433fe8f0b9e 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeUtil.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeUtil.java @@ -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 collectTypeComponentsFromType(@Nullable PyType type, + @NotNull TypeEvalContext context) { + Set 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); + } } diff --git a/python/python-psi-impl/src/com/jetbrains/python/refactoring/PyRefactoringUiService.java b/python/python-psi-impl/src/com/jetbrains/python/refactoring/PyRefactoringUiService.java index 13578f860a04..4b6eb7777f9a 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/refactoring/PyRefactoringUiService.java +++ b/python/python-psi-impl/src/com/jetbrains/python/refactoring/PyRefactoringUiService.java @@ -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 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() { diff --git a/python/python-psi-impl/src/com/jetbrains/python/refactoring/extractmethod/PyExtractMethodSettings.java b/python/python-psi-impl/src/com/jetbrains/python/refactoring/extractmethod/PyExtractMethodSettings.java index 6b68eba77308..3933d43d6f6c 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/refactoring/extractmethod/PyExtractMethodSettings.java +++ b/python/python-psi-impl/src/com/jetbrains/python/refactoring/extractmethod/PyExtractMethodSettings.java @@ -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 { 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 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 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 getReturnTypeFqns() { + return myReturnTypes; + } + public boolean isUseTypeAnnotations() { return myUseTypeAnnotations; } @@ -42,4 +55,15 @@ public class PyExtractMethodSettings implements ExtractMethodSettings { public @Nullable Object getVisibility() { return null; } + + List getAllTypes() { + List result = new ArrayList<>(); + for (PyVariableData variableData : myVariableData) { + if (variableData.type != null) { + result.add(variableData.type); + } + } + result.addAll(myReturnTypes); + return result; + } } diff --git a/python/python-psi-impl/src/com/jetbrains/python/refactoring/extractmethod/PyExtractMethodUtil.java b/python/python-psi-impl/src/com/jetbrains/python/refactoring/extractmethod/PyExtractMethodUtil.java index 44cdcfd54685..822f4c995dd4 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/refactoring/extractmethod/PyExtractMethodUtil.java +++ b/python/python-psi-impl/src/com/jetbrains/python/refactoring/extractmethod/PyExtractMethodUtil.java @@ -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 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 pointers, + @NotNull PyExtractMethodSettings methodSettings, + @NotNull PyFunction insertedMethod, + @NotNull List 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 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; } + } diff --git a/python/python-psi-impl/src/com/jetbrains/python/refactoring/extractmethod/PyVariableData.java b/python/python-psi-impl/src/com/jetbrains/python/refactoring/extractmethod/PyVariableData.java index ea58fb40fcc7..c5ad2325e905 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/refactoring/extractmethod/PyVariableData.java +++ b/python/python-psi-impl/src/com/jetbrains/python/refactoring/extractmethod/PyVariableData.java @@ -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; + } } diff --git a/python/python-psi-impl/src/com/jetbrains/python/refactoring/introduce/IntroduceHandler.java b/python/python-psi-impl/src/com/jetbrains/python/refactoring/introduce/IntroduceHandler.java index 12684e5d7089..7a525ec89b14 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/refactoring/introduce/IntroduceHandler.java +++ b/python/python-psi-impl/src/com/jetbrains/python/refactoring/introduce/IntroduceHandler.java @@ -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; diff --git a/python/src/com/jetbrains/python/refactoring/PyExtractMethodDialog.java b/python/src/com/jetbrains/python/refactoring/PyExtractMethodDialog.java index 705d907afa60..89b2efd13215 100644 --- a/python/src/com/jetbrains/python/refactoring/PyExtractMethodDialog.java +++ b/python/src/com/jetbrains/python/refactoring/PyExtractMethodDialog.java @@ -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 { @@ -62,7 +63,7 @@ public class PyExtractMethodDialog extends AbstractExtractMethodDialog { @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 { 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 { @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 } } diff --git a/python/src/com/jetbrains/python/refactoring/PyRefactoringUiServiceImpl.java b/python/src/com/jetbrains/python/refactoring/PyRefactoringUiServiceImpl.java index e289e8a0ba65..75656619320e 100644 --- a/python/src/com/jetbrains/python/refactoring/PyRefactoringUiServiceImpl.java +++ b/python/src/com/jetbrains/python/refactoring/PyRefactoringUiServiceImpl.java @@ -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(); } diff --git a/python/testData/intentions/PyAnnotateVariableTypeIntentionTest/AnnotationImportTypingAny/main_after.py b/python/testData/intentions/PyAnnotateVariableTypeIntentionTest/AnnotationImportTypingAny/main_after.py index f663ded9e3db..4315c6bf6db0 100644 --- a/python/testData/intentions/PyAnnotateVariableTypeIntentionTest/AnnotationImportTypingAny/main_after.py +++ b/python/testData/intentions/PyAnnotateVariableTypeIntentionTest/AnnotationImportTypingAny/main_after.py @@ -1,4 +1,4 @@ -from typing import Dict, Any +from typing import Any, Dict def func(x): diff --git a/python/testData/intentions/PyAnnotateVariableTypeIntentionTest/AnnotationLiteralEnumType/m.py b/python/testData/intentions/PyAnnotateVariableTypeIntentionTest/AnnotationLiteralEnumType/m.py new file mode 100644 index 000000000000..82a667132a70 --- /dev/null +++ b/python/testData/intentions/PyAnnotateVariableTypeIntentionTest/AnnotationLiteralEnumType/m.py @@ -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]: ... \ No newline at end of file diff --git a/python/testData/intentions/PyAnnotateVariableTypeIntentionTest/AnnotationLiteralEnumType/main.py b/python/testData/intentions/PyAnnotateVariableTypeIntentionTest/AnnotationLiteralEnumType/main.py new file mode 100644 index 000000000000..45198755c3b2 --- /dev/null +++ b/python/testData/intentions/PyAnnotateVariableTypeIntentionTest/AnnotationLiteralEnumType/main.py @@ -0,0 +1,3 @@ +from m import foo_bar + +var = foo_bar() diff --git a/python/testData/intentions/PyAnnotateVariableTypeIntentionTest/AnnotationLiteralEnumType/main_after.py b/python/testData/intentions/PyAnnotateVariableTypeIntentionTest/AnnotationLiteralEnumType/main_after.py new file mode 100644 index 000000000000..d06e713ae51f --- /dev/null +++ b/python/testData/intentions/PyAnnotateVariableTypeIntentionTest/AnnotationLiteralEnumType/main_after.py @@ -0,0 +1,5 @@ +from typing import Literal + +from m import foo_bar, MyEnum + +var: [Literal[MyEnum.A]] = foo_bar() diff --git a/python/testData/intentions/PyAnnotateVariableTypeIntentionTest/annotationGenericParametrizedWithAny_after.py b/python/testData/intentions/PyAnnotateVariableTypeIntentionTest/annotationGenericParametrizedWithAny_after.py index 6f3685b11bf7..519dd1590f91 100644 --- a/python/testData/intentions/PyAnnotateVariableTypeIntentionTest/annotationGenericParametrizedWithAny_after.py +++ b/python/testData/intentions/PyAnnotateVariableTypeIntentionTest/annotationGenericParametrizedWithAny_after.py @@ -1,4 +1,4 @@ -from typing import List, Any +from typing import Any, List def func(): diff --git a/python/testData/intentions/SpecifyTypeInPy3AnnotationsIntentionTest/addsImportsWhenNeeded_after.py b/python/testData/intentions/SpecifyTypeInPy3AnnotationsIntentionTest/addsImportsWhenNeeded_after.py index 701e3699a062..9e08252ab8e1 100644 --- a/python/testData/intentions/SpecifyTypeInPy3AnnotationsIntentionTest/addsImportsWhenNeeded_after.py +++ b/python/testData/intentions/SpecifyTypeInPy3AnnotationsIntentionTest/addsImportsWhenNeeded_after.py @@ -1,4 +1,4 @@ -from typing import Coroutine, Any +from typing import Any, Coroutine async def bar() -> int: diff --git a/python/testData/quickFixes/PyMakeFunctionReturnTypeQuickFixTest/makeGenerator_after.py b/python/testData/quickFixes/PyMakeFunctionReturnTypeQuickFixTest/makeGenerator_after.py index 4f28dff17568..2389737db0b8 100644 --- a/python/testData/quickFixes/PyMakeFunctionReturnTypeQuickFixTest/makeGenerator_after.py +++ b/python/testData/quickFixes/PyMakeFunctionReturnTypeQuickFixTest/makeGenerator_after.py @@ -1,4 +1,4 @@ -from typing import AsyncGenerator, Any +from typing import Any, AsyncGenerator async def gen() -> AsyncGenerator[str | float, Any]: diff --git a/python/testData/refactoring/extractmethod/AsyncDef.after.withTypes.py b/python/testData/refactoring/extractmethod/AsyncDef.after.withTypes.py index e424f420ec81..93e1b6461a6c 100644 --- a/python/testData/refactoring/extractmethod/AsyncDef.after.withTypes.py +++ b/python/testData/refactoring/extractmethod/AsyncDef.after.withTypes.py @@ -1,3 +1,6 @@ +from typing import Any + + async def foo(x): y = await bar(x) return await y diff --git a/python/testData/refactoring/extractmethod/AwaitExpression.after.withTypes.py b/python/testData/refactoring/extractmethod/AwaitExpression.after.withTypes.py index 71c97a3b3c9d..9f6549b6d104 100644 --- a/python/testData/refactoring/extractmethod/AwaitExpression.after.withTypes.py +++ b/python/testData/refactoring/extractmethod/AwaitExpression.after.withTypes.py @@ -1,3 +1,6 @@ +from typing import Any + + async def foo(x): y = await bar(x) return y diff --git a/python/testData/refactoring/extractmethod/BinaryExpression.after.withTypes.py b/python/testData/refactoring/extractmethod/BinaryExpression.after.withTypes.py index f13aaf335d67..9c59f3d363a8 100644 --- a/python/testData/refactoring/extractmethod/BinaryExpression.after.withTypes.py +++ b/python/testData/refactoring/extractmethod/BinaryExpression.after.withTypes.py @@ -1,3 +1,6 @@ +from typing import Any + + def foo() -> Any: return bbb diff --git a/python/testData/refactoring/extractmethod/BreakAst.after.withTypes.py b/python/testData/refactoring/extractmethod/BreakAst.after.withTypes.py index 8184d77a91b4..4884374522e3 100644 --- a/python/testData/refactoring/extractmethod/BreakAst.after.withTypes.py +++ b/python/testData/refactoring/extractmethod/BreakAst.after.withTypes.py @@ -1,3 +1,6 @@ +from typing import Any + + def cylinder_volume(r, h): h * bar(r) diff --git a/python/testData/refactoring/extractmethod/CommentsPrecedingSourceStatement.after.withTypes.py b/python/testData/refactoring/extractmethod/CommentsPrecedingSourceStatement.after.withTypes.py index 878339889f87..d5123b36b4a4 100644 --- a/python/testData/refactoring/extractmethod/CommentsPrecedingSourceStatement.after.withTypes.py +++ b/python/testData/refactoring/extractmethod/CommentsPrecedingSourceStatement.after.withTypes.py @@ -1,3 +1,5 @@ +from typing import Any + x = 42 # print('commented') diff --git a/python/testData/refactoring/extractmethod/ConditionOfConditionalExpression.after.withTypes.py b/python/testData/refactoring/extractmethod/ConditionOfConditionalExpression.after.withTypes.py index 5bdce44c24f0..d4a4a9c14b28 100644 --- a/python/testData/refactoring/extractmethod/ConditionOfConditionalExpression.after.withTypes.py +++ b/python/testData/refactoring/extractmethod/ConditionOfConditionalExpression.after.withTypes.py @@ -1,3 +1,6 @@ +from typing import Any + + def f(n): return n * 2 if bar(n) else n + 1 diff --git a/python/testData/refactoring/extractmethod/DuplicateCheckParam.after.withTypes.py b/python/testData/refactoring/extractmethod/DuplicateCheckParam.after.withTypes.py index ccd298ce79fd..77283869bbc1 100644 --- a/python/testData/refactoring/extractmethod/DuplicateCheckParam.after.withTypes.py +++ b/python/testData/refactoring/extractmethod/DuplicateCheckParam.after.withTypes.py @@ -1,3 +1,6 @@ +from typing import Any + + def f(): a = do_smth() b1 = foo(a) diff --git a/python/testData/refactoring/extractmethod/ElseBody.after.withTypes.py b/python/testData/refactoring/extractmethod/ElseBody.after.withTypes.py index e2012ac11b77..cdeab336ef73 100644 --- a/python/testData/refactoring/extractmethod/ElseBody.after.withTypes.py +++ b/python/testData/refactoring/extractmethod/ElseBody.after.withTypes.py @@ -1,3 +1,7 @@ +from io import TextIOWrapper, _WrappedBuffer +from typing import Any, IO + + def foo(): for arg in sys.argv[1:]: try: diff --git a/python/testData/refactoring/extractmethod/ExtractAddsImport1.after.py b/python/testData/refactoring/extractmethod/ExtractAddsImport1.after.py new file mode 100644 index 000000000000..fbcf844f43e4 --- /dev/null +++ b/python/testData/refactoring/extractmethod/ExtractAddsImport1.after.py @@ -0,0 +1,8 @@ +def f(a): + compiled = compile("x = 42", "", "exec") + body(compiled) + + +def body(compiled_new): + 1 + compiled_new diff --git a/python/testData/refactoring/extractmethod/ExtractAddsImport1.after.withTypes.py b/python/testData/refactoring/extractmethod/ExtractAddsImport1.after.withTypes.py new file mode 100644 index 000000000000..2c5b224e59eb --- /dev/null +++ b/python/testData/refactoring/extractmethod/ExtractAddsImport1.after.withTypes.py @@ -0,0 +1,11 @@ +from types import CodeType + + +def f(a): + compiled = compile("x = 42", "", "exec") + body(compiled) + + +def body(compiled_new: CodeType): + 1 + compiled_new diff --git a/python/testData/refactoring/extractmethod/ExtractAddsImport1.before.py b/python/testData/refactoring/extractmethod/ExtractAddsImport1.before.py new file mode 100644 index 000000000000..66a63f6f438e --- /dev/null +++ b/python/testData/refactoring/extractmethod/ExtractAddsImport1.before.py @@ -0,0 +1,4 @@ +def f(a): + compiled = compile("x = 42", "", "exec") + 1 + compiled diff --git a/python/testData/refactoring/extractmethod/ExtractAddsImport2.after.py b/python/testData/refactoring/extractmethod/ExtractAddsImport2.after.py new file mode 100644 index 000000000000..3bfad4bf034f --- /dev/null +++ b/python/testData/refactoring/extractmethod/ExtractAddsImport2.after.py @@ -0,0 +1,8 @@ +def f(a): + file = open("test.txt", "w") + body(file) + + +def body(file_new): + 1 + file_new diff --git a/python/testData/refactoring/extractmethod/ExtractAddsImport2.after.withTypes.py b/python/testData/refactoring/extractmethod/ExtractAddsImport2.after.withTypes.py new file mode 100644 index 000000000000..98417b789390 --- /dev/null +++ b/python/testData/refactoring/extractmethod/ExtractAddsImport2.after.withTypes.py @@ -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 diff --git a/python/testData/refactoring/extractmethod/ExtractAddsImport2.before.py b/python/testData/refactoring/extractmethod/ExtractAddsImport2.before.py new file mode 100644 index 000000000000..47256fd8f344 --- /dev/null +++ b/python/testData/refactoring/extractmethod/ExtractAddsImport2.before.py @@ -0,0 +1,4 @@ +def f(a): + file = open("test.txt", "w") + 1 + file diff --git a/python/testData/refactoring/extractmethod/ExtractAddsImport3.after.py b/python/testData/refactoring/extractmethod/ExtractAddsImport3.after.py new file mode 100644 index 000000000000..4c4f1aa83e60 --- /dev/null +++ b/python/testData/refactoring/extractmethod/ExtractAddsImport3.after.py @@ -0,0 +1,8 @@ +def f(a): + if a is 1: + body(a) + + +def body(a_new): + 1 + a_new diff --git a/python/testData/refactoring/extractmethod/ExtractAddsImport3.after.withTypes.py b/python/testData/refactoring/extractmethod/ExtractAddsImport3.after.withTypes.py new file mode 100644 index 000000000000..998c8d3ea88e --- /dev/null +++ b/python/testData/refactoring/extractmethod/ExtractAddsImport3.after.withTypes.py @@ -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 diff --git a/python/testData/refactoring/extractmethod/ExtractAddsImport3.before.py b/python/testData/refactoring/extractmethod/ExtractAddsImport3.before.py new file mode 100644 index 000000000000..83d43de134b7 --- /dev/null +++ b/python/testData/refactoring/extractmethod/ExtractAddsImport3.before.py @@ -0,0 +1,4 @@ +def f(a): + if a is 1: + 1 + a diff --git a/python/testData/refactoring/extractmethod/ExtractAddsImport4.after.py b/python/testData/refactoring/extractmethod/ExtractAddsImport4.after.py new file mode 100644 index 000000000000..23176f07e2f9 --- /dev/null +++ b/python/testData/refactoring/extractmethod/ExtractAddsImport4.after.py @@ -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 \ No newline at end of file diff --git a/python/testData/refactoring/extractmethod/ExtractAddsImport4.after.withTypes.py b/python/testData/refactoring/extractmethod/ExtractAddsImport4.after.withTypes.py new file mode 100644 index 000000000000..36096a73639c --- /dev/null +++ b/python/testData/refactoring/extractmethod/ExtractAddsImport4.after.withTypes.py @@ -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 \ No newline at end of file diff --git a/python/testData/refactoring/extractmethod/ExtractAddsImport4.before.py b/python/testData/refactoring/extractmethod/ExtractAddsImport4.before.py new file mode 100644 index 000000000000..487deb1dbfdf --- /dev/null +++ b/python/testData/refactoring/extractmethod/ExtractAddsImport4.before.py @@ -0,0 +1,11 @@ +from enum import Enum + +class Color(Enum): + RED = 1 + GREEN = 2 + BLUE = 3 + +def f(color): + if color == Color.RED: + 1 + color \ No newline at end of file diff --git a/python/testData/refactoring/extractmethod/ExtractAddsImport5.after.py b/python/testData/refactoring/extractmethod/ExtractAddsImport5.after.py new file mode 100644 index 000000000000..6ba36155e717 --- /dev/null +++ b/python/testData/refactoring/extractmethod/ExtractAddsImport5.after.py @@ -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 diff --git a/python/testData/refactoring/extractmethod/ExtractAddsImport5.after.withTypes.py b/python/testData/refactoring/extractmethod/ExtractAddsImport5.after.withTypes.py new file mode 100644 index 000000000000..d450f39e6414 --- /dev/null +++ b/python/testData/refactoring/extractmethod/ExtractAddsImport5.after.withTypes.py @@ -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 diff --git a/python/testData/refactoring/extractmethod/ExtractAddsImport5.before.py b/python/testData/refactoring/extractmethod/ExtractAddsImport5.before.py new file mode 100644 index 000000000000..35e551c316b9 --- /dev/null +++ b/python/testData/refactoring/extractmethod/ExtractAddsImport5.before.py @@ -0,0 +1,7 @@ +import signal + + +def f(sign) : + if sign is signal.Handlers.SIG_DFL: + 1 + sign diff --git a/python/testData/refactoring/extractmethod/IfConditionExpression.after.withTypes.py b/python/testData/refactoring/extractmethod/IfConditionExpression.after.withTypes.py index 54033adc9e5d..76f9af726cd8 100644 --- a/python/testData/refactoring/extractmethod/IfConditionExpression.after.withTypes.py +++ b/python/testData/refactoring/extractmethod/IfConditionExpression.after.withTypes.py @@ -1,3 +1,6 @@ +from typing import Any + + def foo(some_var): if bar(some_var): print('w00t') diff --git a/python/testData/refactoring/extractmethod/IfElseConditionExpression.after.withTypes.py b/python/testData/refactoring/extractmethod/IfElseConditionExpression.after.withTypes.py index e4f5a193a78e..683658654062 100644 --- a/python/testData/refactoring/extractmethod/IfElseConditionExpression.after.withTypes.py +++ b/python/testData/refactoring/extractmethod/IfElseConditionExpression.after.withTypes.py @@ -1,3 +1,6 @@ +from typing import Any + + def foo(some_var): if bar(some_var): print('w00t') diff --git a/python/testData/refactoring/extractmethod/LocalFunction.after.withTypes.py b/python/testData/refactoring/extractmethod/LocalFunction.after.withTypes.py index 32f5eab8f4ce..33bc23ef3e84 100644 --- a/python/testData/refactoring/extractmethod/LocalFunction.after.withTypes.py +++ b/python/testData/refactoring/extractmethod/LocalFunction.after.withTypes.py @@ -1,3 +1,6 @@ +from typing import Any, Callable + + def foo(): def f(x): return x diff --git a/python/testData/refactoring/extractmethod/MethodInnerFuncCombined.after.withTypes.py b/python/testData/refactoring/extractmethod/MethodInnerFuncCombined.after.withTypes.py index e5f563e4bee2..83745c660a9b 100644 --- a/python/testData/refactoring/extractmethod/MethodInnerFuncCombined.after.withTypes.py +++ b/python/testData/refactoring/extractmethod/MethodInnerFuncCombined.after.withTypes.py @@ -1,3 +1,6 @@ +from typing import Any + + class Test: a = 5 def method(self, b): diff --git a/python/testData/refactoring/extractmethod/MethodInnerFuncRecursive.after.withTypes.py b/python/testData/refactoring/extractmethod/MethodInnerFuncRecursive.after.withTypes.py index 5949a0367d5c..cd46fdb7a4fe 100644 --- a/python/testData/refactoring/extractmethod/MethodInnerFuncRecursive.after.withTypes.py +++ b/python/testData/refactoring/extractmethod/MethodInnerFuncRecursive.after.withTypes.py @@ -1,3 +1,6 @@ +from typing import Any + + class Test: def method(self, a): def func(b): diff --git a/python/testData/refactoring/extractmethod/MethodInnerFuncWithMethodParam.after.withTypes.py b/python/testData/refactoring/extractmethod/MethodInnerFuncWithMethodParam.after.withTypes.py index 0df46b87b509..149aa2b45135 100644 --- a/python/testData/refactoring/extractmethod/MethodInnerFuncWithMethodParam.after.withTypes.py +++ b/python/testData/refactoring/extractmethod/MethodInnerFuncWithMethodParam.after.withTypes.py @@ -1,3 +1,6 @@ +from typing import Any + + class Test: def method(self, x): def func(): diff --git a/python/testData/refactoring/extractmethod/MethodInnerFuncWithOwnParam.after.withTypes.py b/python/testData/refactoring/extractmethod/MethodInnerFuncWithOwnParam.after.withTypes.py index 76e57682c3d1..c2e02fbb70f0 100644 --- a/python/testData/refactoring/extractmethod/MethodInnerFuncWithOwnParam.after.withTypes.py +++ b/python/testData/refactoring/extractmethod/MethodInnerFuncWithOwnParam.after.withTypes.py @@ -1,3 +1,6 @@ +from typing import Any + + class Test: def method(self): def func(x): diff --git a/python/testData/refactoring/extractmethod/ReturnTuple.after.withTypes.py b/python/testData/refactoring/extractmethod/ReturnTuple.after.withTypes.py index 1175c453b072..20472802526b 100644 --- a/python/testData/refactoring/extractmethod/ReturnTuple.after.withTypes.py +++ b/python/testData/refactoring/extractmethod/ReturnTuple.after.withTypes.py @@ -1,3 +1,6 @@ +from typing import Any + + def x(p_name, params): return bar(p_name, params), None diff --git a/python/testData/refactoring/extractmethod/SimilarBinaryExpressions.after.withTypes.py b/python/testData/refactoring/extractmethod/SimilarBinaryExpressions.after.withTypes.py index ed03169aa2ab..e0b9bbee2fb4 100644 --- a/python/testData/refactoring/extractmethod/SimilarBinaryExpressions.after.withTypes.py +++ b/python/testData/refactoring/extractmethod/SimilarBinaryExpressions.after.withTypes.py @@ -1,3 +1,6 @@ +from typing import Any + + def compound_duplicate(p1, p2): print(bar(p1)) print(bar(p2)) diff --git a/python/testData/refactoring/extractmethod/StartInArgumentListOfMultiLineFunctionCall.after.withTypes.py b/python/testData/refactoring/extractmethod/StartInArgumentListOfMultiLineFunctionCall.after.withTypes.py index 22351406b71c..4f393cb3eab2 100644 --- a/python/testData/refactoring/extractmethod/StartInArgumentListOfMultiLineFunctionCall.after.withTypes.py +++ b/python/testData/refactoring/extractmethod/StartInArgumentListOfMultiLineFunctionCall.after.withTypes.py @@ -1,3 +1,6 @@ +from typing import Any + + def long_function_name(**kwargs): ... def example_function(): diff --git a/python/testData/refactoring/extractmethod/StartOnMultiLineFunctionCall.after.withTypes.py b/python/testData/refactoring/extractmethod/StartOnMultiLineFunctionCall.after.withTypes.py index 22351406b71c..4f393cb3eab2 100644 --- a/python/testData/refactoring/extractmethod/StartOnMultiLineFunctionCall.after.withTypes.py +++ b/python/testData/refactoring/extractmethod/StartOnMultiLineFunctionCall.after.withTypes.py @@ -1,3 +1,6 @@ +from typing import Any + + def long_function_name(**kwargs): ... def example_function(): diff --git a/python/testData/refactoring/extractmethod/TryContext.after.withTypes.py b/python/testData/refactoring/extractmethod/TryContext.after.withTypes.py index 9378660d87fc..c47b98a8c31c 100644 --- a/python/testData/refactoring/extractmethod/TryContext.after.withTypes.py +++ b/python/testData/refactoring/extractmethod/TryContext.after.withTypes.py @@ -1,3 +1,6 @@ +from typing import Any + + def foo(f): x = 1 x = bar(f, x) diff --git a/python/testData/refactoring/extractmethod/YieldFrom33.after.withTypes.py b/python/testData/refactoring/extractmethod/YieldFrom33.after.withTypes.py index 72bc15729cf1..05add98c8854 100644 --- a/python/testData/refactoring/extractmethod/YieldFrom33.after.withTypes.py +++ b/python/testData/refactoring/extractmethod/YieldFrom33.after.withTypes.py @@ -1,3 +1,6 @@ +from typing import Any + + def f(x, y): yield 'foo' return x, y diff --git a/python/testSrc/com/jetbrains/python/intentions/PyAnnotateVariableTypeIntentionTest.java b/python/testSrc/com/jetbrains/python/intentions/PyAnnotateVariableTypeIntentionTest.java index 7ca94d49c593..ad9cc9f7246c 100644 --- a/python/testSrc/com/jetbrains/python/intentions/PyAnnotateVariableTypeIntentionTest.java +++ b/python/testSrc/com/jetbrains/python/intentions/PyAnnotateVariableTypeIntentionTest.java @@ -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"); diff --git a/python/testSrc/com/jetbrains/python/refactoring/PyExtractMethodTest.java b/python/testSrc/com/jetbrains/python/refactoring/PyExtractMethodTest.java index 3cecdfeb0bed..730f4a2ab93b 100644 --- a/python/testSrc/com/jetbrains/python/refactoring/PyExtractMethodTest.java +++ b/python/testSrc/com/jetbrains/python/refactoring/PyExtractMethodTest.java @@ -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");