correctly handle parenthesises expressions in 'as' clause of with statement (PY-2691)

This commit is contained in:
Dmitry Jemerov
2011-01-14 18:50:38 +01:00
parent 15b2459afd
commit 39b4c5cb24
5 changed files with 33 additions and 12 deletions
@@ -1,10 +1,13 @@
package com.jetbrains.python.psi;
import org.jetbrains.annotations.Nullable;
/**
* @author yole
*/
public interface PyWithItem extends PyElement {
PyWithItem[] EMPTY_ARRAY = new PyWithItem[0];
PyTargetExpression getTargetExpression();
@Nullable
PyExpression getTargetExpression();
}
@@ -1,9 +1,10 @@
package com.jetbrains.python.psi.impl;
import com.intellij.lang.ASTNode;
import com.jetbrains.python.PyElementTypes;
import com.jetbrains.python.psi.PyTargetExpression;
import com.jetbrains.python.PyTokenTypes;
import com.jetbrains.python.psi.PyExpression;
import com.jetbrains.python.psi.PyWithItem;
import org.jetbrains.annotations.Nullable;
/**
* @author yole
@@ -13,9 +14,18 @@ public class PyWithItemImpl extends PyElementImpl implements PyWithItem {
super(astNode);
}
public PyTargetExpression getTargetExpression() {
final ASTNode asNameNode = getNode().findChildByType(PyElementTypes.TARGET_EXPRESSION);
if (asNameNode == null) return null;
return (PyTargetExpression)asNameNode.getPsi();
@Nullable
public PyExpression getTargetExpression() {
ASTNode[] children = getNode().getChildren(null);
boolean foundAs = false;
for (ASTNode child : children) {
if (child.getElementType() == PyTokenTypes.AS_KEYWORD) {
foundAs = true;
}
else if (foundAs && child.getPsi() instanceof PyExpression) {
return (PyExpression) child.getPsi();
}
}
return null;
}
}
@@ -5,10 +5,7 @@ import com.intellij.psi.PsiElement;
import com.intellij.psi.tree.TokenSet;
import com.intellij.psi.util.PsiTreeUtil;
import com.jetbrains.python.PyElementTypes;
import com.jetbrains.python.psi.PyElement;
import com.jetbrains.python.psi.PyElementVisitor;
import com.jetbrains.python.psi.PyWithItem;
import com.jetbrains.python.psi.PyWithStatement;
import com.jetbrains.python.psi.*;
import org.jetbrains.annotations.NotNull;
import java.util.ArrayList;
@@ -34,7 +31,8 @@ public class PyWithStatementImpl extends PyElementImpl implements PyWithStatemen
List<PyElement> result = new ArrayList<PyElement>();
if (items != null) {
for (PyWithItem item : items) {
result.add(item.getTargetExpression());
PyExpression targetExpression = item.getTargetExpression();
result.addAll(PyUtil.flattenedParens(targetExpression));
}
}
return result;
@@ -0,0 +1,6 @@
from contextlib import nested
with nested(patch('Package.ModuleName.ClassName'),
patch('Package.ModuleName.ClassName2', TestUtils.MockClass2)) as (MockClass1, MockClass2):
MockClass1.test.return_value = True
# <ref>
MockClass2.test.return_value = True
@@ -384,4 +384,8 @@ public class PyResolveTest extends PyResolveTestCase {
final PsiElement element = doResolve();
assertNull(element);
}
public void testWithParentheses() {
assertResolvesTo(LanguageLevel.PYTHON27, PyTargetExpression.class, "MockClass1");
}
}