mirror of
https://gitflic.ru/project/openide/openide.git
synced 2026-09-27 10:03:11 +07:00
Type inference for with statements according to context management procotol (PY-4198)
This commit is contained in:
@@ -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()
|
||||
@@ -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();
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user