diff --git a/python/src/com/jetbrains/python/psi/PyUtil.java b/python/src/com/jetbrains/python/psi/PyUtil.java index 6cf072dc6a4a..8c1e318f0bf3 100644 --- a/python/src/com/jetbrains/python/psi/PyUtil.java +++ b/python/src/com/jetbrains/python/psi/PyUtil.java @@ -1718,6 +1718,66 @@ public class PyUtil { return ScratchFileService.isInScratchRoot(PsiUtilCore.getVirtualFile(element)); } + @Nullable + public static PyType getReturnTypeOfMember(@NotNull PyType type, + @NotNull String memberName, + @Nullable PyExpression location, + @NotNull TypeEvalContext context) { + final PyResolveContext resolveContext = PyResolveContext.noImplicits().withTypeEvalContext(context); + final List resolveResults = type.resolveMember(memberName, location, AccessDirection.READ, + resolveContext); + + if (resolveResults != null) { + final List types = new ArrayList<>(); + + for (RatedResolveResult resolveResult : resolveResults) { + final PyType returnType = getReturnType(resolveResult.getElement(), context); + + if (returnType != null) { + types.add(returnType); + } + } + + return PyUnionType.union(types); + } + + return null; + } + + @Nullable + private static PyType getReturnType(@Nullable PsiElement element, @NotNull TypeEvalContext context) { + if (element instanceof PyTypedElement) { + final PyType type = context.getType((PyTypedElement)element); + + return getReturnType(type, context); + } + + return null; + } + + @Nullable + private static PyType getReturnType(@Nullable PyType type, @NotNull TypeEvalContext context) { + if (type instanceof PyCallableType) { + return ((PyCallableType)type).getReturnType(context); + } + + if (type instanceof PyUnionType) { + final List types = new ArrayList<>(); + + for (PyType pyType : ((PyUnionType)type).getMembers()) { + final PyType returnType = getReturnType(pyType, context); + + if (returnType != null) { + types.add(returnType); + } + } + + return PyUnionType.union(types); + } + + return null; + } + /** * This helper class allows to collect various information about AST nodes composing {@link PyStringLiteralExpression}. */ diff --git a/python/src/com/jetbrains/python/psi/impl/PySliceExpressionImpl.java b/python/src/com/jetbrains/python/psi/impl/PySliceExpressionImpl.java index 8a28b9dc5dda..20a27d6ca6bb 100644 --- a/python/src/com/jetbrains/python/psi/impl/PySliceExpressionImpl.java +++ b/python/src/com/jetbrains/python/psi/impl/PySliceExpressionImpl.java @@ -16,21 +16,17 @@ package com.jetbrains.python.psi.impl; import com.intellij.lang.ASTNode; -import com.intellij.psi.PsiElement; import com.intellij.psi.util.PsiTreeUtil; import com.jetbrains.python.PyNames; import com.jetbrains.python.PythonDialectsTokenSetProvider; -import com.jetbrains.python.psi.*; -import com.jetbrains.python.psi.resolve.PyResolveContext; -import com.jetbrains.python.psi.resolve.RatedResolveResult; +import com.jetbrains.python.psi.PyExpression; +import com.jetbrains.python.psi.PySliceExpression; +import com.jetbrains.python.psi.PySliceItem; +import com.jetbrains.python.psi.PyUtil; import com.jetbrains.python.psi.types.*; import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; -import java.util.ArrayList; -import java.util.Collections; -import java.util.List; - /** * @author yole */ @@ -54,24 +50,7 @@ public class PySliceExpressionImpl extends PyElementImpl implements PySliceExpre } if (type instanceof PyClassType) { - final List resolveResults = type.resolveMember( - PyNames.GETITEM, - null, - AccessDirection.READ, - PyResolveContext.noImplicits().withTypeEvalContext(context) - ); - - if (resolveResults != null) { - final List types = new ArrayList<>(); - - for (RatedResolveResult resolveResult : resolveResults) { - types.addAll( - getPossibleReturnTypes(resolveResult.getElement(), context) - ); - } - - return PyUnionType.union(types); - } + return PyUtil.getReturnTypeOfMember(type, PyNames.GETITEM, null, context); } return null; @@ -88,32 +67,4 @@ public class PySliceExpressionImpl extends PyElementImpl implements PySliceExpre public PySliceItem getSliceItem() { return PsiTreeUtil.getChildOfType(this, PySliceItem.class); } - - @NotNull - private static List getPossibleReturnTypes(@Nullable PsiElement element, @NotNull TypeEvalContext context) { - final List result = new ArrayList(); - - if (element instanceof PyTypedElement) { - final PyType elementType = context.getType((PyTypedElement)element); - - result.addAll(getPossibleReturnTypes(elementType, context)); - - if (elementType instanceof PyUnionType) { - for (PyType type : ((PyUnionType)elementType).getMembers()) { - result.addAll(getPossibleReturnTypes(type, context)); - } - } - } - - return result; - } - - @NotNull - private static List getPossibleReturnTypes(@Nullable PyType type, @NotNull TypeEvalContext context) { - if (type instanceof PyCallableType) { - return Collections.singletonList(((PyCallableType)type).getReturnType(context)); - } - - return Collections.emptyList(); - } } diff --git a/python/src/com/jetbrains/python/psi/types/PyClassTypeImpl.java b/python/src/com/jetbrains/python/psi/types/PyClassTypeImpl.java index dcdd2767e6a6..4bba084e6963 100644 --- a/python/src/com/jetbrains/python/psi/types/PyClassTypeImpl.java +++ b/python/src/com/jetbrains/python/psi/types/PyClassTypeImpl.java @@ -380,54 +380,11 @@ public class PyClassTypeImpl extends UserDataHolderBase implements PyClassType { @Nullable private PyType getReturnType(@NotNull TypeEvalContext context, @Nullable PyCallSiteExpression callSite) { if (!isDefinition()) { - final PyResolveContext resolveContext = PyResolveContext.noImplicits().withTypeEvalContext(context); - final List resolveResults = resolveMember(PyNames.CALL, callSite, AccessDirection.READ, resolveContext); - - if (resolveResults != null) { - final ArrayList result = new ArrayList(); - - for (RatedResolveResult resolveResult : resolveResults) { - result.addAll( - getPossibleReturnTypes(resolveResult.getElement(), context) - ); - } - - return PyUnionType.union(result); - } + return PyUtil.getReturnTypeOfMember(this, PyNames.CALL, callSite, context); } else { return new PyClassTypeImpl(getPyClass(), false); } - - return null; - } - - @NotNull - private static List getPossibleReturnTypes(@Nullable PsiElement element, @NotNull TypeEvalContext context) { - final ArrayList result = new ArrayList(); - - if (element instanceof PyTypedElement) { - final PyType elementType = context.getType((PyTypedElement)element); - - result.addAll(getPossibleReturnTypes(elementType, context)); - - if (elementType instanceof PyUnionType) { - for (PyType type : ((PyUnionType)elementType).getMembers()) { - result.addAll(getPossibleReturnTypes(type, context)); - } - } - } - - return result; - } - - @NotNull - private static List getPossibleReturnTypes(@Nullable PyType type, @NotNull TypeEvalContext context) { - if (type instanceof PyCallableType) { - return Collections.singletonList(((PyCallableType)type).getReturnType(context)); - } - - return Collections.emptyList(); } @Nullable