Introduce getIteratedItemType() to PyCollectionType and use it everywhere it's possible

This commit is contained in:
Semyon Proshev
2016-10-13 21:28:20 +03:00
committed by Semyon Proshev
parent 740ef94bcf
commit bfeaf9fe24
13 changed files with 53 additions and 71 deletions
@@ -292,8 +292,8 @@ public class PyStdlibTypeProvider extends PyTypeProviderBase {
return Ref.create(argumentType);
}
else if (argumentType instanceof PyCollectionType) {
final PyType iteratedValueType = ContainerUtil.getFirstItem(((PyCollectionType)argumentType).getElementTypes(context));
return Ref.create(PyTupleType.createHomogeneous(call, iteratedValueType));
final PyType iteratedItemType = ((PyCollectionType)argumentType).getIteratedItemType();
return Ref.create(PyTupleType.createHomogeneous(call, iteratedItemType));
}
return null;
@@ -202,7 +202,7 @@ public class PyTypeModelBuilder {
final PyTupleType tupleType = (PyTupleType)type;
final List<PyType> elementTypes = tupleType.isHomogeneous()
? Collections.singletonList(tupleType.getElementType(0))
? Collections.singletonList(tupleType.getIteratedItemType())
: tupleType.getElementTypes(myContext);
final List<TypeModel> elementModels = ContainerUtil.map(elementTypes, elementType -> build(elementType, true));
@@ -110,7 +110,7 @@ public class PyTypeCheckerInspection extends PyInspection {
final PyType returnType = myTypeEvalContext.getReturnType(function);
if (returnType instanceof PyCollectionType && PyNames.FAKE_COROUTINE.equals(returnType.getName())) {
return ((PyCollectionType)returnType).getElementTypes(myTypeEvalContext).get(0);
return ((PyCollectionType)returnType).getIteratedItemType();
}
return returnType;
@@ -217,7 +217,7 @@ public class PyBuiltinCache {
final List<PyType> tupleElementTypes = tupleType.getElementTypes(context);
if (tupleType.isHomogeneous()) {
final PyType keyAndValueType = tupleElementTypes.get(0);
final PyType keyAndValueType = tupleType.getIteratedItemType();
keyTypes.add(keyAndValueType);
valueTypes.add(keyAndValueType);
@@ -187,7 +187,7 @@ public class PyFunctionImpl extends PyBaseElementImpl<PyFunctionStub> implements
@Override
public PyType getReturnType(@NotNull TypeEvalContext context, @NotNull TypeEvalContext.Key key) {
final PyType type = getReturnType(context);
return isAsync() && isAsyncAllowed() ? createCoroutineType(type, context) : type;
return isAsync() && isAsyncAllowed() ? createCoroutineType(type) : type;
}
@Nullable
@@ -326,14 +326,8 @@ public class PyFunctionImpl extends PyBaseElementImpl<PyFunctionStub> implements
final PyType type = expr != null ? context.getType(expr) : null;
if (node.isDelegating()) {
if (type instanceof PyTupleType) {
types.addAll(((PyTupleType)type).getElementTypes(context));
}
else if (type instanceof PyCollectionType) {
final PyCollectionType collectionType = (PyCollectionType)type;
// TODO: Select the parameter types that matches T in Iterable[T]
final List<PyType> elementTypes = collectionType.getElementTypes(context);
types.add(elementTypes.isEmpty() ? null : elementTypes.get(0));
if (type instanceof PyCollectionType) {
types.add(((PyCollectionType)type).getIteratedItemType());
}
else if (ArrayUtil.contains(type, cache.getListType(), cache.getDictType(), cache.getSetType(), cache.getTupleType())) {
types.add(null);
@@ -387,18 +381,17 @@ public class PyFunctionImpl extends PyBaseElementImpl<PyFunctionStub> implements
}
@Nullable
private PyType createCoroutineType(@Nullable PyType returnType, @NotNull TypeEvalContext context) {
private PyType createCoroutineType(@Nullable PyType returnType) {
final PyBuiltinCache cache = PyBuiltinCache.getInstance(this);
if (returnType instanceof PyCollectionType && PyNames.FAKE_GENERATOR.equals(returnType.getName())) {
final PyClass asyncGenerator = cache.getClass(PyNames.FAKE_ASYNC_GENERATOR);
final List<PyType> generatorElementTypes = ((PyCollectionType)returnType).getElementTypes(context);
if (asyncGenerator == null || generatorElementTypes.isEmpty()) {
if (asyncGenerator == null) {
return null;
}
return new PyCollectionTypeImpl(asyncGenerator, false, Arrays.asList(generatorElementTypes.get(0), null));
return new PyCollectionTypeImpl(asyncGenerator, false, Arrays.asList(((PyCollectionType)returnType).getIteratedItemType(), null));
}
final PyClass coroutine = cache.getClass(PyNames.FAKE_COROUTINE);
@@ -141,32 +141,16 @@ public class PyPrefixExpressionImpl extends PyElementImpl implements PyPrefixExp
@Nullable
private static PyType getGeneratorReturnType(@Nullable PyType type, @NotNull TypeEvalContext context) {
if (type instanceof PyClassLikeType) {
final PyClassLikeType classLikeType = (PyClassLikeType)type;
if (type instanceof PyClassLikeType && type instanceof PyCollectionType) {
// TODO: Understand typing.Generator as well
final String classQName = classLikeType.getClassQName();
final String classQName = ((PyClassLikeType)type).getClassQName();
final PyCollectionType collectionType = (PyCollectionType)type;
if (PyNames.FAKE_GENERATOR.equals(classQName)) {
if (type instanceof PyCollectionType) {
final PyCollectionType collectionType = (PyCollectionType)type;
final List<PyType> elementTypes = collectionType.getElementTypes(context);
if (elementTypes.size() == 3) {
return elementTypes.get(2);
}
}
return ContainerUtil.getOrElse(collectionType.getElementTypes(context), 2, null);
}
else if (PyNames.FAKE_COROUTINE.equals(classQName)) {
if (type instanceof PyCollectionType) {
final PyCollectionType collectionType = (PyCollectionType)type;
final List<PyType> elementTypes = collectionType.getElementTypes(context);
if (elementTypes.size() == 1) {
return elementTypes.get(0);
}
}
}
else if (type instanceof PyClassType &&
type instanceof PyCollectionType &&
PyNames.AWAITABLE.equals(((PyClassType)type).getPyClass().getName())) {
return ContainerUtil.getFirstItem(((PyCollectionType)type).getElementTypes(context));
else if (PyNames.FAKE_COROUTINE.equals(classQName) ||
type instanceof PyClassType && PyNames.AWAITABLE.equals(((PyClassType)type).getPyClass().getName())) {
return collectionType.getIteratedItemType();
}
}
else if (type instanceof PyUnionType) {
@@ -95,9 +95,7 @@ public class PySubscriptionExpressionImpl extends PyElementImpl implements PySub
.orElse(null);
}
else if (type instanceof PyCollectionType) {
// TODO: Select the parameter type that matches T in Iterable[T]
final List<PyType> elementTypes = ((PyCollectionType)type).getElementTypes(context);
res = elementTypes.isEmpty() ? null : elementTypes.get(0);
res = ((PyCollectionType)type).getIteratedItemType();
}
}
}
@@ -346,11 +346,7 @@ public class PyTargetExpressionImpl extends PyBaseElementImpl<PyTargetExpression
@NotNull TypeEvalContext context) {
if (iterableType instanceof PyTupleType) {
final PyTupleType tupleType = (PyTupleType)iterableType;
final List<PyType> memberTypes = new ArrayList<>();
for (int i = 0; i < (tupleType.isHomogeneous() ? 1 : tupleType.getElementCount()); i++) {
memberTypes.add(tupleType.getElementType(i));
}
return PyUnionType.union(memberTypes);
return tupleType.getIteratedItemType();
}
else if (iterableType instanceof PyUnionType) {
final Collection<PyType> members = ((PyUnionType)iterableType).getMembers();
@@ -364,7 +360,7 @@ public class PyTargetExpressionImpl extends PyBaseElementImpl<PyTargetExpression
final PyFunction iterateMethod = findMethodByName(iterableType, PyNames.ITER, context);
if (iterateMethod != null) {
final PyType iterateReturnType = getContextSensitiveType(iterateMethod, context, source);
return getCollectionElementType(iterateReturnType, context);
return getCollectionElementType(iterateReturnType);
}
final String nextMethodName = LanguageLevel.forElement(anchor).isAtLeast(LanguageLevel.PYTHON30) ?
PyNames.DUNDER_NEXT : PyNames.NEXT;
@@ -381,18 +377,16 @@ public class PyTargetExpressionImpl extends PyBaseElementImpl<PyTargetExpression
final PyFunction iterateMethod = findMethodByName(iterableType, PyNames.AITER, context);
if (iterateMethod != null) {
final PyType iterateReturnType = getContextSensitiveType(iterateMethod, context, source);
return getCollectionElementType(iterateReturnType, context);
return getCollectionElementType(iterateReturnType);
}
}
return null;
}
@Nullable
private static PyType getCollectionElementType(@Nullable PyType type, @NotNull TypeEvalContext context) {
private static PyType getCollectionElementType(@Nullable PyType type) {
if (type instanceof PyCollectionType) {
final List<PyType> elementTypes = ((PyCollectionType)type).getElementTypes(context);
// TODO: Select the parameter type that matches T in Iterable[T]
return elementTypes.isEmpty() ? null : elementTypes.get(0);
return ((PyCollectionType)type).getIteratedItemType();
}
return null;
}
@@ -1,5 +1,5 @@
/*
* Copyright 2000-2014 JetBrains s.r.o.
* Copyright 2000-2016 JetBrains s.r.o.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
@@ -16,6 +16,7 @@
package com.jetbrains.python.psi.types;
import org.jetbrains.annotations.NotNull;
import org.jetbrains.annotations.Nullable;
import java.util.List;
@@ -25,4 +26,7 @@ import java.util.List;
public interface PyCollectionType extends PyType {
@NotNull
List<PyType> getElementTypes(@NotNull TypeEvalContext context);
@Nullable
PyType getIteratedItemType();
}
@@ -16,6 +16,7 @@
package com.jetbrains.python.psi.types;
import com.intellij.psi.PsiElement;
import com.intellij.util.containers.ContainerUtil;
import com.jetbrains.python.psi.PyCallSiteExpression;
import com.jetbrains.python.psi.PyClass;
import com.jetbrains.python.psi.PyPsiFacade;
@@ -96,4 +97,11 @@ public class PyCollectionTypeImpl extends PyClassTypeImpl implements PyCollectio
}
return result;
}
@Nullable
@Override
public PyType getIteratedItemType() {
// TODO: Select the parameter type that matches T in Iterable[T]
return ContainerUtil.getFirstItem(myElementTypes);
}
}
@@ -62,7 +62,7 @@ public class PyTupleType extends PyClassTypeImpl implements PyCollectionType {
@NotNull
public String getName() {
if (myHomogeneous) {
return "(" + (getTypeName(myElementTypes.get(0))) + ", ...)";
return "(" + (getTypeName(getIteratedItemType())) + ", ...)";
}
return "(" + StringUtil.join(myElementTypes, PyTupleType::getTypeName, ", ") + ")";
}
@@ -86,7 +86,7 @@ public class PyTupleType extends PyClassTypeImpl implements PyCollectionType {
@Nullable
public PyType getElementType(int index) {
if (myHomogeneous) {
return myElementTypes.get(0);
return getIteratedItemType();
}
if (index >= 0 && index < myElementTypes.size()) {
return myElementTypes.get(index);
@@ -127,4 +127,10 @@ public class PyTupleType extends PyClassTypeImpl implements PyCollectionType {
public List<PyType> getElementTypes(@NotNull TypeEvalContext context) {
return myElementTypes;
}
@Nullable
@Override
public PyType getIteratedItemType() {
return PyUnionType.union(myElementTypes);
}
}
@@ -134,7 +134,7 @@ public class PyTypeChecker {
}
}
else if (superTupleType.isHomogeneous() && !subTupleType.isHomogeneous()) {
final PyType expectedElementType = superTupleType.getElementType(0);
final PyType expectedElementType = superTupleType.getIteratedItemType();
for (int i = 0; i < subTupleType.getElementCount(); i++) {
if (!match(expectedElementType, subTupleType.getElementType(i), context)) {
return false;
@@ -146,7 +146,7 @@ public class PyTypeChecker {
return false;
}
else {
return match(superTupleType.getElementType(0), subTupleType.getElementType(0), context);
return match(superTupleType.getIteratedItemType(), subTupleType.getIteratedItemType(), context);
}
}
else if (expected instanceof PyCollectionType && actual instanceof PyTupleType) {
@@ -154,11 +154,8 @@ public class PyTypeChecker {
return false;
}
final PyTupleType actualTupleType = (PyTupleType)actual;
final PyType superElementType = ContainerUtil.getFirstItem(((PyCollectionType)expected).getElementTypes(context));
final PyType subElementType = actualTupleType.isHomogeneous()
? actualTupleType.getElementType(0)
: PyUnionType.union(actualTupleType.getElementTypes(context));
final PyType superElementType = ((PyCollectionType)expected).getIteratedItemType();
final PyType subElementType = ((PyTupleType)actual).getIteratedItemType();
if (!match(superElementType, subElementType, context, substitutions, recursive)) {
return false;
@@ -371,7 +368,7 @@ public class PyTypeChecker {
final PyClass tupleClass = tupleType.getPyClass();
final List<PyType> oldElementTypes = tupleType.isHomogeneous()
? Collections.singletonList(tupleType.getElementType(0))
? Collections.singletonList(tupleType.getIteratedItemType())
: tupleType.getElementTypes(context);
final List<PyType> newElementTypes =
@@ -54,7 +54,7 @@ public class PyTypeParserTest extends PyTestCase {
final PyCollectionType type = (PyCollectionType) PyTypeParser.getTypeByName(myFixture.getFile(), "list of MyObject");
assertNotNull(type);
assertClassType(type, "list");
assertClassType(type.getElementTypes(getTypeEvalContext()).get(0), "MyObject");
assertClassType(type.getIteratedItemType(), "MyObject");
}
public void testDictType() {
@@ -182,8 +182,7 @@ public class PyTypeParserTest extends PyTestCase {
final PyCollectionType collectionType = (PyCollectionType)type;
assertNotNull(collectionType);
assertEquals("list", collectionType.getName());
final List<PyType> elementTypes = collectionType.getElementTypes(TypeEvalContext.codeInsightFallback(null));
assertInstanceOf(elementTypes.get(0), PyUnionType.class);
assertInstanceOf(collectionType.getIteratedItemType(), PyUnionType.class);
}
public void testBoundedGeneric() {
@@ -203,8 +202,7 @@ public class PyTypeParserTest extends PyTestCase {
final PyCollectionType collectionType = (PyCollectionType)type;
assertNotNull(collectionType);
assertEquals("list", collectionType.getName());
final List<PyType> elementTypes = collectionType.getElementTypes(TypeEvalContext.codeInsightFallback(null));
assertEquals("int", elementTypes.get(0).getName());
assertEquals("int", collectionType.getIteratedItemType().getName());
}
public void testBracketMultipleParams() {