diff --git a/python/src/com/jetbrains/python/codeInsight/stdlib/PyStdlibTypeProvider.java b/python/src/com/jetbrains/python/codeInsight/stdlib/PyStdlibTypeProvider.java index 5125a710f906..6d3d6164f87b 100644 --- a/python/src/com/jetbrains/python/codeInsight/stdlib/PyStdlibTypeProvider.java +++ b/python/src/com/jetbrains/python/codeInsight/stdlib/PyStdlibTypeProvider.java @@ -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; diff --git a/python/src/com/jetbrains/python/documentation/PyTypeModelBuilder.java b/python/src/com/jetbrains/python/documentation/PyTypeModelBuilder.java index 6c2ba27fdff6..1bbea0644db6 100644 --- a/python/src/com/jetbrains/python/documentation/PyTypeModelBuilder.java +++ b/python/src/com/jetbrains/python/documentation/PyTypeModelBuilder.java @@ -202,7 +202,7 @@ public class PyTypeModelBuilder { final PyTupleType tupleType = (PyTupleType)type; final List elementTypes = tupleType.isHomogeneous() - ? Collections.singletonList(tupleType.getElementType(0)) + ? Collections.singletonList(tupleType.getIteratedItemType()) : tupleType.getElementTypes(myContext); final List elementModels = ContainerUtil.map(elementTypes, elementType -> build(elementType, true)); diff --git a/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java b/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java index cc22e31b92dc..0ff3c667f448 100644 --- a/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java +++ b/python/src/com/jetbrains/python/inspections/PyTypeCheckerInspection.java @@ -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; diff --git a/python/src/com/jetbrains/python/psi/impl/PyBuiltinCache.java b/python/src/com/jetbrains/python/psi/impl/PyBuiltinCache.java index 81d2eede1826..8ce29a200eda 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyBuiltinCache.java +++ b/python/src/com/jetbrains/python/psi/impl/PyBuiltinCache.java @@ -217,7 +217,7 @@ public class PyBuiltinCache { final List tupleElementTypes = tupleType.getElementTypes(context); if (tupleType.isHomogeneous()) { - final PyType keyAndValueType = tupleElementTypes.get(0); + final PyType keyAndValueType = tupleType.getIteratedItemType(); keyTypes.add(keyAndValueType); valueTypes.add(keyAndValueType); diff --git a/python/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java b/python/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java index 3697fdb32ead..44e45ed23a39 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java +++ b/python/src/com/jetbrains/python/psi/impl/PyFunctionImpl.java @@ -187,7 +187,7 @@ public class PyFunctionImpl extends PyBaseElementImpl 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 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 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 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 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); diff --git a/python/src/com/jetbrains/python/psi/impl/PyPrefixExpressionImpl.java b/python/src/com/jetbrains/python/psi/impl/PyPrefixExpressionImpl.java index 03db867a6942..41173c0526b6 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyPrefixExpressionImpl.java +++ b/python/src/com/jetbrains/python/psi/impl/PyPrefixExpressionImpl.java @@ -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 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 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) { diff --git a/python/src/com/jetbrains/python/psi/impl/PySubscriptionExpressionImpl.java b/python/src/com/jetbrains/python/psi/impl/PySubscriptionExpressionImpl.java index 12817d6345a2..f79369369230 100644 --- a/python/src/com/jetbrains/python/psi/impl/PySubscriptionExpressionImpl.java +++ b/python/src/com/jetbrains/python/psi/impl/PySubscriptionExpressionImpl.java @@ -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 elementTypes = ((PyCollectionType)type).getElementTypes(context); - res = elementTypes.isEmpty() ? null : elementTypes.get(0); + res = ((PyCollectionType)type).getIteratedItemType(); } } } diff --git a/python/src/com/jetbrains/python/psi/impl/PyTargetExpressionImpl.java b/python/src/com/jetbrains/python/psi/impl/PyTargetExpressionImpl.java index 3112396e4302..bb63c93f83ba 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyTargetExpressionImpl.java +++ b/python/src/com/jetbrains/python/psi/impl/PyTargetExpressionImpl.java @@ -346,11 +346,7 @@ public class PyTargetExpressionImpl extends PyBaseElementImpl 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 members = ((PyUnionType)iterableType).getMembers(); @@ -364,7 +360,7 @@ public class PyTargetExpressionImpl extends PyBaseElementImpl 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; } diff --git a/python/src/com/jetbrains/python/psi/types/PyCollectionType.java b/python/src/com/jetbrains/python/psi/types/PyCollectionType.java index 3da104dee6fe..bc9d0370ed36 100644 --- a/python/src/com/jetbrains/python/psi/types/PyCollectionType.java +++ b/python/src/com/jetbrains/python/psi/types/PyCollectionType.java @@ -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 getElementTypes(@NotNull TypeEvalContext context); + + @Nullable + PyType getIteratedItemType(); } diff --git a/python/src/com/jetbrains/python/psi/types/PyCollectionTypeImpl.java b/python/src/com/jetbrains/python/psi/types/PyCollectionTypeImpl.java index e5f2689f7ee2..92bb5a904d1f 100644 --- a/python/src/com/jetbrains/python/psi/types/PyCollectionTypeImpl.java +++ b/python/src/com/jetbrains/python/psi/types/PyCollectionTypeImpl.java @@ -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); + } } \ No newline at end of file diff --git a/python/src/com/jetbrains/python/psi/types/PyTupleType.java b/python/src/com/jetbrains/python/psi/types/PyTupleType.java index 5a98be568e63..c74bdf28ec3d 100644 --- a/python/src/com/jetbrains/python/psi/types/PyTupleType.java +++ b/python/src/com/jetbrains/python/psi/types/PyTupleType.java @@ -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 getElementTypes(@NotNull TypeEvalContext context) { return myElementTypes; } + + @Nullable + @Override + public PyType getIteratedItemType() { + return PyUnionType.union(myElementTypes); + } } diff --git a/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java b/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java index 3d30c0c1ef6e..8d4f3255182f 100644 --- a/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java +++ b/python/src/com/jetbrains/python/psi/types/PyTypeChecker.java @@ -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 oldElementTypes = tupleType.isHomogeneous() - ? Collections.singletonList(tupleType.getElementType(0)) + ? Collections.singletonList(tupleType.getIteratedItemType()) : tupleType.getElementTypes(context); final List newElementTypes = diff --git a/python/testSrc/com/jetbrains/python/PyTypeParserTest.java b/python/testSrc/com/jetbrains/python/PyTypeParserTest.java index bd671c7aa8dd..b0c91362ee1f 100644 --- a/python/testSrc/com/jetbrains/python/PyTypeParserTest.java +++ b/python/testSrc/com/jetbrains/python/PyTypeParserTest.java @@ -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 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 elementTypes = collectionType.getElementTypes(TypeEvalContext.codeInsightFallback(null)); - assertEquals("int", elementTypes.get(0).getName()); + assertEquals("int", collectionType.getIteratedItemType().getName()); } public void testBracketMultipleParams() {