mirror of
https://gitflic.ru/project/openide/openide.git
synced 2026-09-27 10:03:11 +07:00
Fixed type inference for subscription operator of a union type (PY-12862)
This commit is contained in:
@@ -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();
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user