diff --git a/python/src/com/jetbrains/python/refactoring/classes/PyClassRefactoringUtil.java b/python/src/com/jetbrains/python/refactoring/classes/PyClassRefactoringUtil.java index 906a2ad592ec..41ae56c6562c 100644 --- a/python/src/com/jetbrains/python/refactoring/classes/PyClassRefactoringUtil.java +++ b/python/src/com/jetbrains/python/refactoring/classes/PyClassRefactoringUtil.java @@ -28,6 +28,8 @@ import com.intellij.psi.util.PsiTreeUtil; import com.intellij.psi.util.PsiUtilBase; import com.intellij.psi.util.QualifiedName; import com.intellij.util.ArrayUtil; +import com.intellij.util.Function; +import com.intellij.util.containers.ContainerUtil; import com.jetbrains.NotNullPredicate; import com.jetbrains.python.PyNames; import com.jetbrains.python.codeInsight.PyCodeInsightSettings; @@ -51,6 +53,13 @@ public final class PyClassRefactoringUtil { private static final Key ENCODED_IMPORT = Key.create("PyEncodedImport"); private static final Key ENCODED_USE_FROM_IMPORT = Key.create("PyEncodedUseFromImport"); private static final Key ENCODED_IMPORT_AS = Key.create("PyEncodedImportAs"); + /** + * Existence of value for this key means that given element was imported from the module via {@code __all__}, but nonetheless unresolved, + * because it was defined indirectly, e.g. via {@code globals()), {@code sys.modules}, etc. In this case {@link #ENCODED_IMPORT} + * will point to {@code __all__} itself, and the only way to know element's name is to extract it from the corresponding import element + * and then save as copyable user data using this key. + */ + private static final Key DYNAMIC_IMPORTED_NAME = Key.create("PyDynamicImportedName"); private PyClassRefactoringUtil() { @@ -246,32 +255,48 @@ public final class PyClassRefactoringUtil { } - private static void restoreReference(final PyReferenceExpression node, PsiElement[] otherMovedElements) { - PsiNamedElement target = node.getCopyableUserData(ENCODED_IMPORT); - final String asName = node.getCopyableUserData(ENCODED_IMPORT_AS); - final Boolean useFromImport = node.getCopyableUserData(ENCODED_USE_FROM_IMPORT); - if (target instanceof PsiDirectory) { - target = (PsiNamedElement)PyUtil.getPackageElement((PsiDirectory)target, node); - } - if (target instanceof PyFunction) { - final PyFunction f = (PyFunction)target; - final PyClass c = f.getContainingClass(); - if (c != null && c.findInitOrNew(false) == f) { - target = c; + private static void restoreReference(@NotNull PyReferenceExpression node, @NotNull PsiElement[] otherMovedElements) { + try { + PsiNamedElement target = node.getCopyableUserData(ENCODED_IMPORT); + final String asName = node.getCopyableUserData(ENCODED_IMPORT_AS); + final Boolean useFromImport = node.getCopyableUserData(ENCODED_USE_FROM_IMPORT); + if (target instanceof PsiDirectory) { + target = (PsiNamedElement)PyUtil.getPackageElement((PsiDirectory)target, node); + } + if (target instanceof PyFunction) { + final PyFunction f = (PyFunction)target; + final PyClass c = f.getContainingClass(); + if (c != null && c.findInitOrNew(false) == f) { + target = c; + } + } + if (target instanceof PyTargetExpression && PyNames.ALL.equals(target.getName())) { + final PsiFile srcModule = target.getContainingFile(); + final PsiFile destModule = node.getContainingFile(); + final QualifiedName allModulePath = QualifiedNameFinder.findCanonicalImportPath(srcModule, node); + final String importedName = node.getCopyableUserData(DYNAMIC_IMPORTED_NAME); + if (allModulePath != null && importedName != null) { + final AddImportHelper.ImportPriority priority = AddImportHelper.getImportPriority(destModule, srcModule); + AddImportHelper.addOrUpdateFromImportStatement(destModule, allModulePath.toString(), importedName, asName, priority, null); + } + return; + } + if (target == null) return; + if (PsiTreeUtil.isAncestor(node.getContainingFile(), target, false)) return; + if (ArrayUtil.contains(target, otherMovedElements)) return; + if (target instanceof PyFile || target instanceof PsiDirectory) { + insertImport(node, target, asName, useFromImport != null ? useFromImport : true); + } + else { + insertImport(node, target, asName, true); } } - if (target == null) return; - if (PsiTreeUtil.isAncestor(node.getContainingFile(), target, false)) return; - if (ArrayUtil.contains(target, otherMovedElements)) return; - if (target instanceof PyFile || target instanceof PsiDirectory) { - insertImport(node, target, asName, useFromImport != null ? useFromImport : true); + finally { + node.putCopyableUserData(ENCODED_IMPORT, null); + node.putCopyableUserData(ENCODED_IMPORT_AS, null); + node.putCopyableUserData(ENCODED_USE_FROM_IMPORT, null); + node.putCopyableUserData(DYNAMIC_IMPORTED_NAME, null); } - else { - insertImport(node, target, asName, true); - } - node.putCopyableUserData(ENCODED_IMPORT, null); - node.putCopyableUserData(ENCODED_IMPORT_AS, null); - node.putCopyableUserData(ENCODED_USE_FROM_IMPORT, null); } public static void insertImport(PsiElement anchor, Collection elements) { @@ -371,12 +396,24 @@ public final class PyClassRefactoringUtil { if (qualifier != null && !(resolveExpression(qualifier) instanceof PyImportedModule)) { return; } - final PsiElement target = resolveExpression(node); + final List allResolveResults = multiResolveExpression(node); + final PsiElement target = ContainerUtil.getFirstItem(allResolveResults); if (target instanceof PsiNamedElement && !PsiTreeUtil.isAncestor(element, target, false)) { final NameDefiner importElement = getImportElement(node); if (!PyUtil.inSameFile(element, target) && importElement == null && !(target instanceof PsiFileSystemItem)) { return; } + if (target instanceof PyTargetExpression && PyNames.ALL.equals(((PyTargetExpression)target).getName())) { + for (PsiElement result : allResolveResults) { + if (result instanceof PyImportElement) { + final QualifiedName importedQName = ((PyImportElement)result).getImportedQName(); + if (importedQName != null) { + node.putCopyableUserData(DYNAMIC_IMPORTED_NAME, importedQName.toString()); + break; + } + } + } + } node.putCopyableUserData(ENCODED_IMPORT, (PsiNamedElement)target); if (importElement instanceof PyImportElement) { node.putCopyableUserData(ENCODED_IMPORT_AS, ((PyImportElement)importElement).getAsName()); @@ -407,6 +444,16 @@ public final class PyClassRefactoringUtil { return null; } + @NotNull + private static List multiResolveExpression(@NotNull PyReferenceExpression expr) { + return ContainerUtil.mapNotNull(expr.getReference().multiResolve(false), new Function() { + @Override + public PsiElement fun(ResolveResult result) { + return result.getElement(); + } + }); + } + public static void updateImportOfElement(@NotNull PyImportStatementBase importStatement, @NotNull PsiNamedElement element) { final String name = getOriginalName(element); if (name != null) { diff --git a/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAll/after/src/a.py b/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAll/after/src/a.py new file mode 100644 index 000000000000..b28b04f64312 --- /dev/null +++ b/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAll/after/src/a.py @@ -0,0 +1,3 @@ + + + diff --git a/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAll/after/src/b.py b/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAll/after/src/b.py new file mode 100644 index 000000000000..2214497f5839 --- /dev/null +++ b/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAll/after/src/b.py @@ -0,0 +1,18 @@ +__all__ = ['foo'] + + +class C(object): + def foo(self): + pass + + +_c = C() + + +def _init(g): + for name in dir(_c): + if not name.startswith('_'): + g[name] = getattr(_c, name) + + +_init(globals()) \ No newline at end of file diff --git a/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAll/after/src/c.py b/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAll/after/src/c.py new file mode 100644 index 000000000000..9d4cc491d1f2 --- /dev/null +++ b/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAll/after/src/c.py @@ -0,0 +1,5 @@ +from b import foo + + +def use_foo(): + print(foo) \ No newline at end of file diff --git a/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAll/before/src/a.py b/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAll/before/src/a.py new file mode 100644 index 000000000000..9d4cc491d1f2 --- /dev/null +++ b/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAll/before/src/a.py @@ -0,0 +1,5 @@ +from b import foo + + +def use_foo(): + print(foo) \ No newline at end of file diff --git a/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAll/before/src/b.py b/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAll/before/src/b.py new file mode 100644 index 000000000000..2214497f5839 --- /dev/null +++ b/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAll/before/src/b.py @@ -0,0 +1,18 @@ +__all__ = ['foo'] + + +class C(object): + def foo(self): + pass + + +_c = C() + + +def _init(g): + for name in dir(_c): + if not name.startswith('_'): + g[name] = getattr(_c, name) + + +_init(globals()) \ No newline at end of file diff --git a/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAll/before/src/c.py b/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAll/before/src/c.py new file mode 100644 index 000000000000..e69de29bb2d1 diff --git a/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAllWithAlias/after/src/a.py b/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAllWithAlias/after/src/a.py new file mode 100644 index 000000000000..b28b04f64312 --- /dev/null +++ b/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAllWithAlias/after/src/a.py @@ -0,0 +1,3 @@ + + + diff --git a/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAllWithAlias/after/src/b.py b/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAllWithAlias/after/src/b.py new file mode 100644 index 000000000000..2214497f5839 --- /dev/null +++ b/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAllWithAlias/after/src/b.py @@ -0,0 +1,18 @@ +__all__ = ['foo'] + + +class C(object): + def foo(self): + pass + + +_c = C() + + +def _init(g): + for name in dir(_c): + if not name.startswith('_'): + g[name] = getattr(_c, name) + + +_init(globals()) \ No newline at end of file diff --git a/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAllWithAlias/after/src/c.py b/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAllWithAlias/after/src/c.py new file mode 100644 index 000000000000..bbe2c1d10347 --- /dev/null +++ b/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAllWithAlias/after/src/c.py @@ -0,0 +1,5 @@ +from b import foo as bar + + +def use_foo(): + print(bar) \ No newline at end of file diff --git a/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAllWithAlias/before/src/a.py b/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAllWithAlias/before/src/a.py new file mode 100644 index 000000000000..bbe2c1d10347 --- /dev/null +++ b/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAllWithAlias/before/src/a.py @@ -0,0 +1,5 @@ +from b import foo as bar + + +def use_foo(): + print(bar) \ No newline at end of file diff --git a/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAllWithAlias/before/src/b.py b/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAllWithAlias/before/src/b.py new file mode 100644 index 000000000000..2214497f5839 --- /dev/null +++ b/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAllWithAlias/before/src/b.py @@ -0,0 +1,18 @@ +__all__ = ['foo'] + + +class C(object): + def foo(self): + pass + + +_c = C() + + +def _init(g): + for name in dir(_c): + if not name.startswith('_'): + g[name] = getattr(_c, name) + + +_init(globals()) \ No newline at end of file diff --git a/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAllWithAlias/before/src/c.py b/python/testData/refactoring/move/usageFromFunctionResolvesToDunderAllWithAlias/before/src/c.py new file mode 100644 index 000000000000..e69de29bb2d1 diff --git a/python/testSrc/com/jetbrains/python/refactoring/PyMoveTest.java b/python/testSrc/com/jetbrains/python/refactoring/PyMoveTest.java index 22c496fc85f2..5bd85a97663f 100644 --- a/python/testSrc/com/jetbrains/python/refactoring/PyMoveTest.java +++ b/python/testSrc/com/jetbrains/python/refactoring/PyMoveTest.java @@ -354,6 +354,16 @@ public class PyMoveTest extends PyTestCase { doMoveSymbolsTest("b.py", "func", "C"); } + // PY-14811 + public void testUsageFromFunctionResolvesToDunderAll() { + doMoveSymbolTest("use_foo", "c.py"); + } + + // PY-14811 + public void testUsageFromFunctionResolvesToDunderAllWithAlias() { + doMoveSymbolTest("use_foo", "c.py"); + } + private void doMoveFileTest(String fileName, String toDirName) { Project project = myFixture.getProject(); PsiManager manager = PsiManager.getInstance(project);