Type inference for with statements according to context management procotol (PY-4198)

This commit is contained in:
Andrey Vlasovskikh
2011-07-27 22:39:49 +04:00
parent 8338191ded
commit ee7d005612
8 changed files with 74 additions and 8 deletions
@@ -78,6 +78,7 @@ public class PyNames {
public static final String GETITEM = "__getitem__";
public static final String NAME = "__name__";
public static final String ENTER = "__enter__";
/**
* Contains all known predefined names of "__foo__" form.
@@ -9,5 +9,8 @@ public interface PyWithItem extends PyElement {
PyWithItem[] EMPTY_ARRAY = new PyWithItem[0];
@Nullable
PyExpression getTargetExpression();
PyExpression getExpression();
@Nullable
PyExpression getTarget();
}
@@ -136,8 +136,9 @@ public class PyTargetExpressionImpl extends PyPresentableElementImpl<PyTargetExp
if (type != null) {
return type;
}
if (getParent() instanceof PyAssignmentStatement) {
final PyAssignmentStatement assignmentStatement = (PyAssignmentStatement)getParent();
final PsiElement parent = getParent();
if (parent instanceof PyAssignmentStatement) {
final PyAssignmentStatement assignmentStatement = (PyAssignmentStatement)parent;
final PyExpression assignedValue = assignmentStatement.getAssignedValue();
if (assignedValue != null) {
if (assignedValue instanceof PyReferenceExpressionImpl) {
@@ -166,12 +167,26 @@ public class PyTargetExpressionImpl extends PyPresentableElementImpl<PyTargetExp
return context.getType(assignedValue);
}
}
if (getParent() instanceof PyTupleExpression) {
if (parent instanceof PyTupleExpression) {
final PyType typeFromTupleAssignment = getTypeFromTupleAssignment(context);
if (typeFromTupleAssignment != null) {
return typeFromTupleAssignment;
}
}
if (parent instanceof PyWithItem) {
final PyWithItem item = (PyWithItem)parent;
final PyType exprType = item.getExpression().getType(context);
if (exprType instanceof PyClassType) {
final PyClass cls = ((PyClassType)exprType).getPyClass();
if (cls != null) {
final PyFunction enter = cls.findMethodByName(PyNames.ENTER, true);
if (enter != null) {
return enter.getReturnType(context, null);
}
}
}
return null;
}
PyType iterType = getTypeFromIteration(context);
if (iterType != null) {
return iterType;
@@ -1,6 +1,7 @@
package com.jetbrains.python.psi.impl;
import com.intellij.lang.ASTNode;
import com.intellij.psi.PsiElement;
import com.jetbrains.python.PyTokenTypes;
import com.jetbrains.python.psi.PyExpression;
import com.jetbrains.python.psi.PyWithItem;
@@ -14,8 +15,23 @@ public class PyWithItemImpl extends PyElementImpl implements PyWithItem {
super(astNode);
}
@Nullable
public PyExpression getTargetExpression() {
@Override
public PyExpression getExpression() {
ASTNode[] children = getNode().getChildren(null);
for (ASTNode child: children) {
final PsiElement e = child.getPsi();
if (e instanceof PyExpression) {
return (PyExpression)e;
}
else if (child.getElementType() == PyTokenTypes.AS_KEYWORD) {
break;
}
}
return null;
}
@Override
public PyExpression getTarget() {
ASTNode[] children = getNode().getChildren(null);
boolean foundAs = false;
for (ASTNode child : children) {
@@ -31,7 +31,7 @@ public class PyWithStatementImpl extends PyElementImpl implements PyWithStatemen
List<PyElement> result = new ArrayList<PyElement>();
if (items != null) {
for (PyWithItem item : items) {
PyExpression targetExpression = item.getTargetExpression();
PyExpression targetExpression = item.getTarget();
result.addAll(PyUtil.flattenedParensAndTuples(targetExpression));
}
}
@@ -0,0 +1,13 @@
class Eggs(object):
def __enter__(self):
return u'foo'
def __exit__(self, exc_type, exc_val, exc_tb):
pass
class Spam(Eggs):
pass
def f():
with Spam() as spam:
spam.encode()
+13
View File
@@ -0,0 +1,13 @@
class Eggs(object):
def __enter__(self):
return u'foo'
def __exit__(self, exc_type, exc_val, exc_tb):
pass
class Spam(Eggs):
pass
def f():
with Spam() as spam:
spam.enc<caret>
@@ -177,7 +177,12 @@ public class PythonCompletionTest extends PyLightFixtureTestCase {
doTest();
}
public void testReturnType() {
public void testReturnType() {
doTest();
}
public void testWithType() { // PY-4198
setLanguageLevel(LanguageLevel.PYTHON26);
doTest();
}