Fixed type inference for subscription operator of a union type (PY-12862)

This commit is contained in:
Andrey Vlasovskikh
2014-10-15 18:26:03 +04:00
parent 8d07983e50
commit e7f5877888
3 changed files with 66 additions and 26 deletions
@@ -30,6 +30,9 @@ import com.jetbrains.python.psi.types.*;
import org.jetbrains.annotations.NotNull;
import org.jetbrains.annotations.Nullable;
import java.util.ArrayList;
import java.util.List;
/**
* @author yole
*/
@@ -66,29 +69,32 @@ public class PySubscriptionExpressionImpl extends PyElementImpl implements PySub
@Nullable
@Override
public PyType getType(@NotNull TypeEvalContext context, @NotNull TypeEvalContext.Key key) {
PyType res = null;
final PsiReference ref = getReference(PyResolveContext.noImplicits().withTypeEvalContext(context));
final PsiElement resolved = ref.resolve();
if (resolved instanceof Callable) {
res = ((Callable)resolved).getCallType(context, this);
}
if (PyTypeChecker.isUnknown(res) || res instanceof PyNoneType) {
final PyExpression indexExpression = getIndexExpression();
if (indexExpression != null) {
final PyType type = context.getType(getOperand());
final PyClass cls = (type instanceof PyClassType) ? ((PyClassType)type).getPyClass() : null;
if (cls != null && PyABCUtil.isSubclass(cls, PyNames.MAPPING)) {
return res;
}
if (type instanceof PySubscriptableType) {
res = ((PySubscriptableType)type).getElementType(indexExpression, context);
}
else if (type instanceof PyCollectionType) {
res = ((PyCollectionType)type).getElementType(context);
final PsiPolyVariantReference reference = getReference(PyResolveContext.noImplicits().withTypeEvalContext(context));
final List<PyType> members = new ArrayList<PyType>();
for (PsiElement resolved : PyUtil.multiResolveTopPriority(reference)) {
PyType res = null;
if (resolved instanceof Callable) {
res = ((Callable)resolved).getCallType(context, this);
}
if (PyTypeChecker.isUnknown(res) || res instanceof PyNoneType) {
final PyExpression indexExpression = getIndexExpression();
if (indexExpression != null) {
final PyType type = context.getType(getOperand());
final PyClass cls = (type instanceof PyClassType) ? ((PyClassType)type).getPyClass() : null;
if (cls != null && PyABCUtil.isSubclass(cls, PyNames.MAPPING)) {
return res;
}
if (type instanceof PySubscriptableType) {
res = ((PySubscriptableType)type).getElementType(indexExpression, context);
}
else if (type instanceof PyCollectionType) {
res = ((PyCollectionType)type).getElementType(context);
}
}
}
members.add(res);
}
return res;
return PyUnionType.union(members);
}
@Override
@@ -1,5 +1,5 @@
def test():
xs = map(lambda x: x + 1, [1, 2, 3])
print('foo' + <warning descr="Expected type 'str | unicode', got 'int' instead">xs[0]</warning>)
ys = map(str, iter([1, 2, 3]))
print(1 + <warning descr="Expected type 'Number', got 'str' instead">ys[0]</warning>, 'bar' + ys[1])
print('foo' + xs[0]) # Can be a str since map returns list[V] | str | unicode
ys = map(tuple, iter([1, 2, 3]))
print(1 + <warning descr="Expected type 'Number', got 'tuple | str | unicode' instead">ys[0]</warning>, 'bar' + ys[1])
@@ -621,12 +621,25 @@ public class PyTypeTest extends PyTestCase {
}
public void testFunctionTypeAsUnificationArgument() {
doTest("int",
doTest("list[int] | str | unicode",
"def map2(f, xs):\n" +
" '''\n" +
" :type f: (T) -> V | None\n" +
" :type xs: collections.Iterable[T] | bytes | unicode\n" +
" :rtype: list[V] | bytes | unicode\n" +
" :type xs: collections.Iterable[T] | str | unicode\n" +
" :rtype: list[V] | str | unicode\n" +
" '''\n" +
" pass\n" +
"\n" +
"expr = map2(lambda x: 10, ['1', '2', '3'])\n");
}
public void testFunctionTypeAsUnificationArgumentWithSubscription() {
doTest("int | str | unicode",
"def map2(f, xs):\n" +
" '''\n" +
" :type f: (T) -> V | None\n" +
" :type xs: collections.Iterable[T] | str | unicode\n" +
" :rtype: list[V] | str | unicode\n" +
" '''\n" +
" pass\n" +
"\n" +
@@ -906,6 +919,27 @@ public class PyTypeTest extends PyTestCase {
"expr = f().foo()\n");
}
// PY-12862
public void testUnionTypeAttributeSubscriptionOfDifferentTypes() {
doTest("C1 | C2",
"class C1:\n" +
" def __getitem__(self, item):\n" +
" return self\n" +
"\n" +
"class C2:\n" +
" def __getitem__(self, item):\n" +
" return self\n" +
"\n" +
"def f():\n" +
" '''\n" +
" :rtype: C1 | C2\n" +
" '''\n" +
" pass\n" +
"\n" +
"expr = f()[0]\n" +
"print(expr)\n");
}
private static TypeEvalContext getTypeEvalContext(@NotNull PyExpression element) {
return TypeEvalContext.userInitiated(element.getContainingFile()).withTracing();
}