From 32bca31f70c6fcaf86f1c3d3933183faf1d27c04 Mon Sep 17 00:00:00 2001 From: Dmitry Jemerov Date: Thu, 26 Apr 2012 11:21:12 +0200 Subject: [PATCH] fix SOE in new-style resolve (PY-6305) (oh god the contract of findExportedName() is so confusing) --- .../src/com/jetbrains/python/psi/impl/PyFileImpl.java | 10 ++++++---- .../fromPackageImportIntoInit/pack/__init__.py | 2 ++ .../multiFile/fromPackageImportIntoInit/pack/mod.py | 0 .../com/jetbrains/python/PyMultiFileResolveTest.java | 8 ++++++++ 4 files changed, 16 insertions(+), 4 deletions(-) create mode 100644 python/testData/resolve/multiFile/fromPackageImportIntoInit/pack/__init__.py create mode 100644 python/testData/resolve/multiFile/fromPackageImportIntoInit/pack/mod.py diff --git a/python/src/com/jetbrains/python/psi/impl/PyFileImpl.java b/python/src/com/jetbrains/python/psi/impl/PyFileImpl.java index f7bcb9668701..35512c774fac 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyFileImpl.java +++ b/python/src/com/jetbrains/python/psi/impl/PyFileImpl.java @@ -183,9 +183,11 @@ public class PyFileImpl extends PsiFileBase implements PyFile, PyExpression { return null; } + @Nullable private PsiElement resolveDeclaration(String name, PsiElement result) { if (result instanceof PyImportElement) { - return findNameInImportElement(name, (PyImportElement)result); + final PyImportElement importElement = (PyImportElement)result; + return findNameInImportElement(name, importElement, importElement.getContainingImportStatement() instanceof PyFromImportStatement); } else if (result instanceof PyFromImportStatement) { return ((PyFromImportStatement) result).resolveImportSource(); @@ -557,7 +559,7 @@ public class PyFileImpl extends PsiFileBase implements PyFile, PyExpression { @Nullable private PsiElement findNameInImportStatement(String name, PyImportStatement child) { for (PyImportElement importElement: child.getImportElements()) { - final PsiElement result = findNameInImportElement(name, importElement); + final PsiElement result = findNameInImportElement(name, importElement, false); if (result != null) { return result; } @@ -566,8 +568,8 @@ public class PyFileImpl extends PsiFileBase implements PyFile, PyExpression { } @Nullable - private PsiElement findNameInImportElement(String name, PyImportElement importElement) { - final PsiElement result = importElement.getElementNamed(name, false); + private PsiElement findNameInImportElement(String name, PyImportElement importElement, final boolean resolveImportElement) { + final PsiElement result = importElement.getElementNamed(name, resolveImportElement); if (result != null) { return result; } diff --git a/python/testData/resolve/multiFile/fromPackageImportIntoInit/pack/__init__.py b/python/testData/resolve/multiFile/fromPackageImportIntoInit/pack/__init__.py new file mode 100644 index 000000000000..d84254aa0698 --- /dev/null +++ b/python/testData/resolve/multiFile/fromPackageImportIntoInit/pack/__init__.py @@ -0,0 +1,2 @@ +from pack import mod +# diff --git a/python/testData/resolve/multiFile/fromPackageImportIntoInit/pack/mod.py b/python/testData/resolve/multiFile/fromPackageImportIntoInit/pack/mod.py new file mode 100644 index 000000000000..e69de29bb2d1 diff --git a/python/testSrc/com/jetbrains/python/PyMultiFileResolveTest.java b/python/testSrc/com/jetbrains/python/PyMultiFileResolveTest.java index 4540eb1a1480..250977bd7e2f 100644 --- a/python/testSrc/com/jetbrains/python/PyMultiFileResolveTest.java +++ b/python/testSrc/com/jetbrains/python/PyMultiFileResolveTest.java @@ -96,6 +96,14 @@ public class PyMultiFileResolveTest extends PyResolveTestCase { assertTrue("is target?", elt instanceof PyTargetExpression); } + public void testFromPackageImportIntoInit() { // PY-6305 + myFixture.copyDirectoryToProject("fromPackageImportIntoInit/pack", "pack"); + final PsiFile psiFile = myFixture.configureByFile("pack/__init__.py"); + final PsiElement result = doResolve(psiFile); + assertInstanceOf(result, PyFile.class); + assertEquals("mod.py", ((PyFile) result).getName()); + } + public void testResolveInPkg() { ResolveResult[] results = doMultiResolve(); assertTrue(results.length == 2); // func and import stmt