Merge branch 'python-fixes'

This commit is contained in:
Andrey Vlasovskikh
2015-11-16 19:33:47 +03:00
7 changed files with 174 additions and 183 deletions
@@ -15,11 +15,9 @@
*/
package com.jetbrains.python.psi;
import com.intellij.psi.PsiNameIdentifierOwner;
import com.intellij.psi.PsiNamedElement;
import com.intellij.psi.PsiReference;
import com.intellij.psi.StubBasedPsiElement;
import com.intellij.psi.*;
import com.intellij.psi.util.QualifiedName;
import com.jetbrains.python.psi.resolve.PyResolveContext;
import com.jetbrains.python.psi.stubs.PyTargetExpressionStub;
import org.jetbrains.annotations.NotNull;
import org.jetbrains.annotations.Nullable;
@@ -36,11 +34,26 @@ public interface PyTargetExpression extends PyQualifiedExpression, PsiNamedEleme
* Find the value that maps to this target expression in an enclosing assignment expression.
* Does not work with other expressions (e.g. if the target is in a 'for' loop).
*
* Operates at the AST level.
*
* @return the expression assigned to target via an enclosing assignment expression, or null.
*/
@Nullable
PyExpression findAssignedValue();
/**
* Resolves the value that maps to this target expression in an enclosing assignment expression.
*
* This method does not access AST if underlying PSI is stub based and the context doesn't allow switching to AST.
*/
@Nullable
PsiElement resolveAssignedValue(@NotNull PyResolveContext resolveContext);
/**
* Returns the qualified name (if there is any) assigned to the expression.
*
* This method does not access AST if underlying PSI is stub based.
*/
@Nullable
QualifiedName getAssignedQName();
@@ -22,6 +22,7 @@ import com.intellij.psi.PsiElement;
import com.jetbrains.python.PyBundle;
import com.jetbrains.python.psi.*;
import com.jetbrains.python.psi.resolve.PyResolveContext;
import com.jetbrains.python.psi.types.TypeEvalContext;
import org.jetbrains.annotations.NotNull;
public class PyReplaceTupleWithListQuickFix implements LocalQuickFix {
@@ -46,7 +47,9 @@ public class PyReplaceTupleWithListQuickFix implements LocalQuickFix {
PySubscriptionExpression subscriptionExpression = (PySubscriptionExpression)targets[0];
if (subscriptionExpression.getOperand() instanceof PyReferenceExpression) {
PyReferenceExpression referenceExpression = (PyReferenceExpression)subscriptionExpression.getOperand();
element = referenceExpression.followAssignmentsChain(PyResolveContext.defaultContext()).getElement();
final TypeEvalContext context = TypeEvalContext.userInitiated(project, element.getContainingFile());
final PyResolveContext resolveContext = PyResolveContext.noImplicits().withTypeEvalContext(context);
element = referenceExpression.followAssignmentsChain(resolveContext).getElement();
if (element instanceof PyParenthesizedExpression) {
final PyExpression expression = ((PyParenthesizedExpression)element).getContainedExpression();
replaceWithListLiteral(element, (PyTupleExpression)expression);
@@ -60,7 +60,10 @@ public class PyReferenceExpressionImpl extends PyElementImpl implements PyRefere
@NotNull
@Override
public PsiPolyVariantReference getReference() {
return getReference(PyResolveContext.defaultContext());
//noinspection InstanceofIncompatibleInterface
assert !(this instanceof StubBasedPsiElement);
final TypeEvalContext context = TypeEvalContext.codeAnalysis(getProject(), getContainingFile());
return getReference(PyResolveContext.defaultContext().withTypeEvalContext(context));
}
@NotNull
@@ -144,15 +147,14 @@ public class PyReferenceExpressionImpl extends PyElementImpl implements PyRefere
for (ResolveResult target : targets) {
PsiElement elt = target.getElement();
if (elt instanceof PyTargetExpression) {
PsiElement assigned_from = null;
final PyTargetExpression expr = (PyTargetExpression)elt;
final TypeEvalContext context = resolveContext.getTypeEvalContext();
if (context.maySwitchToAST(expr) || expr.getStub() == null) {
final PsiElement assigned_from;
if (context.maySwitchToAST(expr)) {
assigned_from = expr.findAssignedValue();
}
// TODO: Maybe findAssignedValueByStub() should become a part of the PyTargetExpression interface
else if (elt instanceof PyTargetExpressionImpl) {
assigned_from = ((PyTargetExpressionImpl)elt).findAssignedValueByStub(context);
else {
assigned_from = expr.resolveAssignedValue(resolveContext);
}
if (assigned_from instanceof PyReferenceExpression) {
if (visited.contains(assigned_from)) {
@@ -144,17 +144,13 @@ public class PyTargetExpressionImpl extends PyBaseElementImpl<PyTargetExpression
return type;
}
if (!context.maySwitchToAST(this)) {
final PsiElement value = getStub() != null ? findAssignedValueByStub(context) : findAssignedValue();
final PsiElement value = resolveAssignedValue(PyResolveContext.noImplicits().withTypeEvalContext(context));
if (value instanceof PyTypedElement) {
type = context.getType((PyTypedElement)value);
if (type instanceof PyNoneType) {
return null;
}
if (type instanceof PyFunctionTypeImpl) {
return type;
}
// We are unsure about the type since it may be inferred from the stub based on incomplete information
return PyUnionType.createWeakType(type);
return type;
}
return null;
}
@@ -476,6 +472,50 @@ public class PyTargetExpressionImpl extends PyBaseElementImpl<PyTargetExpression
}
@Nullable
@Override
public PsiElement resolveAssignedValue(@NotNull PyResolveContext resolveContext) {
final TypeEvalContext context = resolveContext.getTypeEvalContext();
if (context.maySwitchToAST(this)) {
final PyExpression value = findAssignedValue();
if (value != null) {
final List<PsiElement> results = PyUtil.multiResolveTopPriority(value, resolveContext);
return !results.isEmpty() ? results.get(0) : null;
}
return null;
}
else {
final QualifiedName qName = getAssignedQName();
if (qName != null) {
final ScopeOwner owner = ScopeUtil.getScopeOwner(this);
if (owner instanceof PyTypedElement) {
final List<String> components = qName.getComponents();
if (!components.isEmpty()) {
PsiElement resolved = owner;
for (String component : components) {
if (!(resolved instanceof PyTypedElement)) {
return null;
}
final PyType qualifierType = context.getType((PyTypedElement)resolved);
if (qualifierType == null) {
return null;
}
final List<? extends RatedResolveResult> results = qualifierType.resolveMember(component, null, AccessDirection.READ,
resolveContext);
if (results == null || results.isEmpty()) {
return null;
}
resolved = results.get(0).getElement();
}
return resolved;
}
}
}
return null;
}
}
@Nullable
@Override
public PyExpression findAssignedValue() {
if (isValid()) {
PyAssignmentStatement assignment = PsiTreeUtil.getParentOfType(this, PyAssignmentStatement.class);
@@ -490,6 +530,8 @@ public class PyTargetExpressionImpl extends PyBaseElementImpl<PyTargetExpression
return null;
}
@Nullable
@Override
public QualifiedName getAssignedQName() {
final PyTargetExpressionStub stub = getStub();
if (stub != null) {
@@ -501,35 +543,6 @@ public class PyTargetExpressionImpl extends PyBaseElementImpl<PyTargetExpression
return PyPsiUtils.asQualifiedName(findAssignedValue());
}
@Nullable
public PsiElement findAssignedValueByStub(@NotNull TypeEvalContext context) {
final PyTargetExpressionStub stub = getStub();
if (stub != null && stub.getInitializerType() == PyTargetExpressionStub.InitializerType.ReferenceExpression) {
final QualifiedName initializer = stub.getInitializer();
// TODO: Support qualified stub initializers
if (initializer != null && initializer.getComponentCount() == 1) {
final String name = initializer.getLastComponent();
if (name != null) {
final PsiElement parent = getParentByStub();
if (parent instanceof PyFile) {
return ((PyFile)parent).getElementNamed(name);
}
else if (parent instanceof PyClass) {
final PyType type = context.getType((PyClass)parent);
if (type != null) {
final List<? extends RatedResolveResult> results = type.resolveMember(name, null, AccessDirection.READ,
PyResolveContext.noImplicits());
if (results != null && !results.isEmpty()) {
return results.get(0).getElement();
}
}
}
}
}
}
return null;
}
@Override
public QualifiedName getCalleeName() {
final PyTargetExpressionStub stub = getStub();
@@ -15,6 +15,8 @@
*/
package com.jetbrains.python;
import com.intellij.openapi.project.Project;
import com.intellij.psi.PsiFile;
import com.intellij.testFramework.LightProjectDescriptor;
import com.jetbrains.python.documentation.PythonDocumentationProvider;
import com.jetbrains.python.fixtures.PyTestCase;
@@ -164,13 +166,17 @@ public class Py3TypeTest extends PyTestCase {
}
});
}
private void doTest(final String expectedType, final String text) {
myFixture.configureByText(PythonFileType.INSTANCE, text);
final PyExpression expr = myFixture.findElementByText("expr", PyExpression.class);
final TypeEvalContext context = TypeEvalContext.userInitiated(expr.getProject(), expr.getContainingFile()).withTracing();
final Project project = expr.getProject();
final PsiFile containingFile = expr.getContainingFile();
assertType(expectedType, expr, TypeEvalContext.codeAnalysis(project, containingFile));
assertType(expectedType, expr, TypeEvalContext.userInitiated(project, containingFile));
}
private static void assertType(String expectedType, PyExpression expr, TypeEvalContext context) {
final PyType actual = context.getType(expr);
final String actualType = PythonDocumentationProvider.getTypeName(actual, context);
assertEquals(expectedType, actualType);
@@ -21,6 +21,7 @@ import com.jetbrains.python.fixtures.PyTestCase;
import com.jetbrains.python.psi.PyCallExpression;
import com.jetbrains.python.psi.PyFunction;
import com.jetbrains.python.psi.resolve.PyResolveContext;
import com.jetbrains.python.psi.types.TypeEvalContext;
/**
* Tests callee resolution in PyCallExpressionImpl.
@@ -32,7 +33,8 @@ public class PyResolveCalleeTest extends PyTestCase {
private PyCallExpression.PyMarkedCallee resolveCallee() {
PsiReference ref = myFixture.getReferenceAtCaretPosition("/resolve/callee/" + getTestName(false) + ".py");
PyCallExpression call = PsiTreeUtil.getParentOfType(ref.getElement(), PyCallExpression.class);
return call.resolveCallee(PyResolveContext.defaultContext());
final TypeEvalContext context = TypeEvalContext.codeAnalysis(myFixture.getProject(), myFixture.getFile());
return call.resolveCallee(PyResolveContext.noImplicits().withTypeEvalContext(context));
}
public void testInstanceCall() {
@@ -15,13 +15,17 @@
*/
package com.jetbrains.python;
import com.google.common.collect.ImmutableList;
import com.jetbrains.python.documentation.PythonDocumentationProvider;
import com.jetbrains.python.fixtures.PyTestCase;
import com.jetbrains.python.psi.LanguageLevel;
import com.jetbrains.python.psi.PyExpression;
import com.jetbrains.python.psi.impl.PythonLanguageLevelPusher;
import com.jetbrains.python.psi.types.*;
import com.jetbrains.python.psi.types.PyClassType;
import com.jetbrains.python.psi.types.PyType;
import com.jetbrains.python.psi.types.TypeEvalContext;
import org.jetbrains.annotations.NotNull;
import org.jetbrains.annotations.Nullable;
import java.util.List;
@@ -134,6 +138,7 @@ public class PyTypeTest extends PyTestCase {
// TODO: uncomment when we have a mock SDK for Python 3.x
// PY-1427
@SuppressWarnings("unused")
public void _testBytesLiteral() { // PY-1427
PythonLanguageLevelPusher.setForcedLanguageLevel(myFixture.getProject(), LanguageLevel.PYTHON30);
try {
@@ -321,59 +326,42 @@ public class PyTypeTest extends PyTestCase {
}
public void testSOEOnRecursiveCall() {
PyExpression expr = parseExpr("def foo(x): return foo(x)\n" +
"expr = foo(1)");
TypeEvalContext context = getTypeEvalContext(expr);
PyType actual = context.getType(expr);
assertNull(actual);
doTest("Any", "def foo(x): return foo(x)\n" +
"expr = foo(1)");
}
public void testGenericConcrete() {
PyExpression expr = parseExpr("def f(x):\n" +
" '''\n" +
" :type x: T\n" +
" :rtype: T\n" +
" '''\n" +
" return x\n" +
"\n" +
"expr = f(1)\n");
TypeEvalContext context = getTypeEvalContext(expr);
PyType actual = context.getType(expr);
assertNotNull(actual);
assertEquals("int", actual.getName());
doTest("int", "def f(x):\n" +
" '''\n" +
" :type x: T\n" +
" :rtype: T\n" +
" '''\n" +
" return x\n" +
"\n" +
"expr = f(1)\n");
}
public void testGenericConcreteMismatch() {
PyExpression expr = parseExpr("def f(x, y):\n" +
" '''\n" +
" :type x: T\n" +
" :rtype: T\n" +
" '''\n" +
" return x\n" +
"\n" +
"expr = f(1)\n");
TypeEvalContext context = getTypeEvalContext(expr);
PyType actual = context.getType(expr);
assertNotNull(actual);
assertEquals("int", actual.getName());
doTest("int", "def f(x, y):\n" +
" '''\n" +
" :type x: T\n" +
" :rtype: T\n" +
" '''\n" +
" return x\n" +
"\n" +
"expr = f(1)\n");
}
// PY-5831
public void testYieldType() {
PyExpression expr = parseExpr("def f():\n" +
" expr = yield 2\n");
TypeEvalContext context = getTypeEvalContext(expr);
PyType actual = context.getType(expr);
assertNull(actual);
doTest("Any", "def f():\n" +
" expr = yield 2\n");
}
// PY-9590
public void testYieldParensType() {
PyExpression expr = parseExpr("def f():\n" +
" expr = (yield 2)\n");
TypeEvalContext context = getTypeEvalContext(expr);
PyType actual = context.getType(expr);
assertNull(actual);
doTest("Any", "def f():\n" +
" expr = (yield 2)\n");
}
public void testFunctionAssignment() {
@@ -386,27 +374,19 @@ public class PyTypeTest extends PyTestCase {
}
public void testPropertyOfUnionType() {
PyExpression expr = parseExpr("def f():\n" +
" '''\n" +
" :rtype: int or slice\n" +
" '''\n" +
" raise NotImplementedError\n" +
"\n" +
"x = f()\n" +
"expr = x.start\n");
TypeEvalContext context = getTypeEvalContext(expr);
PyType actual = context.getType(expr);
assertNotNull(actual);
assertInstanceOf(actual, PyClassType.class);
assertEquals("int", actual.getName());
doTest("int", "def f():\n" +
" '''\n" +
" :rtype: int or slice\n" +
" '''\n" +
" raise NotImplementedError\n" +
"\n" +
"x = f()\n" +
"expr = x.start\n");
}
public void testUndefinedPropertyOfUnionType() {
PyExpression expr = parseExpr("x = 42 if True else 'spam'\n" +
"expr = x.foo\n");
TypeEvalContext context = getTypeEvalContext(expr);
PyType actual = context.getType(expr);
assertNull(actual);
doTest("Any", "x = 42 if True else 'spam'\n" +
"expr = x.foo\n");
}
// PY-7058
@@ -416,31 +396,26 @@ public class PyTypeTest extends PyTestCase {
"\n" +
"x = C()\n" +
"expr = type(x)\n");
TypeEvalContext context = getTypeEvalContext(expr);
PyType type = context.getType(expr);
assertInstanceOf(type, PyClassType.class);
assertTrue("Got instance type instead of class type", ((PyClassType)type).isDefinition());
assertNotNull(expr);
for (TypeEvalContext context : getTypeEvalContexts(expr)) {
PyType type = context.getType(expr);
assertInstanceOf(type, PyClassType.class);
assertTrue("Got instance type instead of class type", ((PyClassType)type).isDefinition());
}
}
// PY-7058
public void testReturnTypeOfTypeForClass() {
PyExpression expr = parseExpr("class C(object):\n" +
" pass\n" +
"\n" +
"expr = type(C)\n");
TypeEvalContext context = getTypeEvalContext(expr);
PyType type = context.getType(expr);
assertInstanceOf(type, PyClassType.class);
assertEquals(type.getName(), "type");
doTest("type", "class C(object):\n" +
" pass\n" +
"\n" +
"expr = type(C)\n");
}
// PY-7058
public void testReturnTypeOfTypeForUnknown() {
PyExpression expr = parseExpr("def f(x):\n" +
" expr = type(x)\n");
TypeEvalContext context = getTypeEvalContext(expr);
PyType type = context.getType(expr);
assertNull(type);
doTest("Any", "def f(x):\n" +
" expr = type(x)\n");
}
// PY-7040
@@ -483,30 +458,12 @@ public class PyTypeTest extends PyTestCase {
// PY-7020
public void testListComprehensionType() {
final PyExpression expr = parseExpr("expr = [str(x) for x in range(10)]\n");
final TypeEvalContext context = getTypeEvalContext(expr);
final PyType type = context.getType(expr);
assertNotNull(type);
assertInstanceOf(type, PyCollectionType.class);
assertEquals("list", type.getName());
final PyCollectionType collectionType = (PyCollectionType)type;
final List<PyType> elementTypes = collectionType.getElementTypes(context);
assertEquals("str", elementTypes.get(0).getName());
doTest("List[str]", "expr = [str(x) for x in range(10)]\n");
}
// PY-7021
public void testGeneratorComprehensionType() {
final PyExpression expr = parseExpr("expr = (str(x) for x in range(10))\n");
final TypeEvalContext context = getTypeEvalContext(expr);
final PyType type = context.getType(expr);
assertNotNull(type);
assertInstanceOf(type, PyCollectionType.class);
assertEquals("__generator", type.getName());
final PyCollectionType collectionType = (PyCollectionType)type;
final List<PyType> elementTypes = collectionType.getElementTypes(context);
assertEquals("str", elementTypes.get(0).getName());
assertTrue(PyTypeChecker.isUnknown(elementTypes.get(1)));
assertEquals("None", elementTypes.get(2).getName());
doTest("__generator[str, Any, None]", "expr = (str(x) for x in range(10))\n");
}
// PY-7021
@@ -595,11 +552,8 @@ public class PyTypeTest extends PyTestCase {
// PY-7063
public void testDefaultParameterIgnoreNone() {
final PyExpression expr = parseExpr("def f(x=None):\n" +
" expr = x\n");
final TypeEvalContext context = getTypeEvalContext(expr);
final PyType type = context.getType(expr);
assertNull(type);
doTest("Any", "def f(x=None):\n" +
" expr = x\n");
}
public void testParameterFromUsages() {
@@ -610,6 +564,7 @@ public class PyTypeTest extends PyTestCase {
" foo(3)\n" +
" foo('bar')\n";
final PyExpression expr = parseExpr(text);
assertNotNull(expr);
doTest("Union[Union[int, str], Any]", expr, TypeEvalContext.codeCompletion(expr.getProject(), expr.getContainingFile()));
}
@@ -675,24 +630,18 @@ public class PyTypeTest extends PyTestCase {
}
public void testUnionIteration() {
final String text = "def f(c):\n" +
" if c < 0:\n" +
" return [1, 2, 3]\n" +
" elif c == 0:\n" +
" return 0.0\n" +
" else:\n" +
" return 'foo'\n" +
"\n" +
"def g(c):\n" +
" for expr in f(c):\n" +
" pass\n";
final PyExpression expr = parseExpr(text);
final TypeEvalContext context = getTypeEvalContext(expr);
final PyType type = context.getType(expr);
assertInstanceOf(type, PyUnionType.class);
assertTrue(PyTypeChecker.match(PyTypeParser.getTypeByName(expr, "int"), type, context));
assertTrue(PyTypeChecker.match(PyTypeParser.getTypeByName(expr, "str"), type, context));
assertTrue(PyTypeChecker.isUnknown(type));
doTest("Union[Union[int, str], Any]",
"def f(c):\n" +
" if c < 0:\n" +
" return [1, 2, 3]\n" +
" elif c == 0:\n" +
" return 0.0\n" +
" else:\n" +
" return 'foo'\n" +
"\n" +
"def g(c):\n" +
" for expr in f(c):\n" +
" pass\n");
}
public void testParameterOfFunctionTypeAndReturnValue() {
@@ -1011,11 +960,13 @@ public class PyTypeTest extends PyTestCase {
"expr = f\n");
}
private static TypeEvalContext getTypeEvalContext(@NotNull PyExpression element) {
return TypeEvalContext.userInitiated(element.getProject(), element.getContainingFile()).withTracing();
private static List<TypeEvalContext> getTypeEvalContexts(@NotNull PyExpression element) {
return ImmutableList.of(TypeEvalContext.codeAnalysis(element.getProject(), element.getContainingFile()).withTracing(),
TypeEvalContext.userInitiated(element.getProject(), element.getContainingFile()).withTracing());
}
private PyExpression parseExpr(String text) {
@Nullable
private PyExpression parseExpr(@NotNull String text) {
myFixture.configureByText(PythonFileType.INSTANCE, text);
return myFixture.findElementByText("expr", PyExpression.class);
}
@@ -1026,22 +977,23 @@ public class PyTypeTest extends PyTestCase {
assertEquals(expectedType, actualType);
}
private void doTest(final String expectedType, final String text) {
PyExpression expr = parseExpr(text);
TypeEvalContext context = getTypeEvalContext(expr);
PyType actual = context.getType(expr);
final String actualType = PythonDocumentationProvider.getTypeName(actual, context);
assertEquals(expectedType, actualType);
private void doTest(@NotNull final String expectedType, @NotNull final String text) {
checkTypes(expectedType, parseExpr(text));
}
private static void checkTypes(@NotNull String expectedType, @Nullable PyExpression expr) {
assertNotNull(expr);
for (TypeEvalContext context : getTypeEvalContexts(expr)) {
final PyType actual = context.getType(expr);
final String actualType = PythonDocumentationProvider.getTypeName(actual, context);
assertEquals("Failed in " + context, expectedType, actualType);
}
}
public static final String TEST_DIRECTORY = "/types/";
private void doMultiFileTest(final String expectedType, final String text) {
private void doMultiFileTest(@NotNull final String expectedType, @NotNull final String text) {
myFixture.copyDirectoryToProject(TEST_DIRECTORY + getTestName(false), "");
PyExpression expr = parseExpr(text);
TypeEvalContext context = getTypeEvalContext(expr);
PyType actual = context.getType(expr);
final String actualType = PythonDocumentationProvider.getTypeName(actual, context);
assertEquals(expectedType, actualType);
checkTypes(expectedType, parseExpr(text));
}
}