diff --git a/python/src/com/jetbrains/python/codeInsight/PyTypingTypeProvider.java b/python/src/com/jetbrains/python/codeInsight/PyTypingTypeProvider.java index f37048a1c21d..d414e9abead1 100644 --- a/python/src/com/jetbrains/python/codeInsight/PyTypingTypeProvider.java +++ b/python/src/com/jetbrains/python/codeInsight/PyTypingTypeProvider.java @@ -83,7 +83,7 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { // XXX: Requires switching from stub to AST final PyExpression value = annotation.getValue(); if (value != null) { - final PyType type = getType(value, context, new HashSet<>()); + final PyType type = getType(value, new Context(context)); if (type != null) { final PyType optionalType = getOptionalTypeFromDefaultNone(param, type, context); return Ref.create(optionalType != null ? optionalType : type); @@ -127,11 +127,11 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { // XXX: Requires switching from stub to AST final PyExpression value = annotation.getValue(); if (value != null) { - final PyType type = getType(value, context, new HashSet<>()); + final PyType type = getType(value, new Context(context)); return type != null ? Ref.create(type) : null; } } - final PyType constructorType = getGenericConstructorType(function, context, new HashSet<>()); + final PyType constructorType = getGenericConstructorType(function, new Context(context)); if (constructorType != null) { return Ref.create(constructorType); } @@ -155,7 +155,7 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { final PyExpression[] args = callExpr.getArguments(); if (args.length > 0) { final PyExpression typeExpr = args[0]; - return getType(typeExpr, context, new HashSet<>()); + return getType(typeExpr, new Context(context)); } } return null; @@ -167,7 +167,7 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { final PyTargetExpression target = (PyTargetExpression)referenceTarget; final String comment = getTypeComment(target); if (comment != null) { - final PyType type = getStringBasedType(comment, referenceTarget, context, new HashSet<>()); + final PyType type = getStringBasedType(comment, referenceTarget, new Context(context)); if (type instanceof PyTupleType) { final PyTupleExpression tupleExpr = PsiTreeUtil.getParentOfType(target, PyTupleExpression.class); if (tupleExpr != null) { @@ -244,11 +244,11 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { } @Nullable - private static PyType getGenericConstructorType(@NotNull PyFunction function, @NotNull TypeEvalContext context, @NotNull Set cache) { + private static PyType getGenericConstructorType(@NotNull PyFunction function, @NotNull Context context) { if (PyUtil.isInit(function)) { final PyClass cls = function.getContainingClass(); if (cls != null) { - final List genericTypes = collectGenericTypes(cls, context, cache); + final List genericTypes = collectGenericTypes(cls, context); final List elementTypes = new ArrayList(genericTypes); if (!elementTypes.isEmpty()) { return new PyCollectionTypeImpl(cls, false, elementTypes); @@ -259,9 +259,9 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { } @NotNull - private static List collectGenericTypes(@NotNull PyClass cls, @NotNull TypeEvalContext context, @NotNull Set cache) { + private static List collectGenericTypes(@NotNull PyClass cls, @NotNull Context context) { boolean isGeneric = false; - for (PyClass ancestor : cls.getAncestorClasses(context)) { + for (PyClass ancestor : cls.getAncestorClasses(context.getTypeContext())) { if (GENERIC_CLASSES.contains(ancestor.getQualifiedName())) { isGeneric = true; break; @@ -274,8 +274,8 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { if (expr instanceof PySubscriptionExpression) { final PyExpression indexExpr = ((PySubscriptionExpression)expr).getIndexExpression(); if (indexExpr != null) { - for (PsiElement resolved : tryResolving(indexExpr, context)) { - final PyGenericType genericType = getGenericType(resolved, context, cache); + for (PsiElement resolved : tryResolving(indexExpr, context.getTypeContext())) { + final PyGenericType genericType = getGenericType(resolved, context); if (genericType != null) { results.add(genericType); } @@ -289,38 +289,36 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { } @Nullable - private static PyType getType(@NotNull PyExpression expression, @NotNull TypeEvalContext context, @NotNull Set cache) { + private static PyType getType(@NotNull PyExpression expression, @NotNull Context context) { final List members = Lists.newArrayList(); - for (PsiElement resolved : tryResolving(expression, context)) { - members.add(getTypeForResolvedElement(resolved, context, cache)); + for (PsiElement resolved : tryResolving(expression, context.getTypeContext())) { + members.add(getTypeForResolvedElement(resolved, context)); } return PyUnionType.union(members); } @Nullable - private static PyType getTypeForResolvedElement(@NotNull PsiElement resolved, - @NotNull TypeEvalContext context, - @NotNull Set cache) { - if (cache.contains(resolved)) { + private static PyType getTypeForResolvedElement(@NotNull PsiElement resolved, @NotNull Context context) { + if (context.getExpressionCache().contains(resolved)) { // Recursive types are not yet supported return null; } - cache.add(resolved); + context.getExpressionCache().add(resolved); try { - final PyType unionType = getUnionType(resolved, context, cache); + final PyType unionType = getUnionType(resolved, context); if (unionType != null) { return unionType; } - final Ref optionalType = getOptionalType(resolved, context, cache); + final Ref optionalType = getOptionalType(resolved, context); if (optionalType != null) { return optionalType.get(); } - final PyType callableType = getCallableType(resolved, context, cache); + final PyType callableType = getCallableType(resolved, context); if (callableType != null) { return callableType; } - final PyType parameterizedType = getParameterizedType(resolved, context, cache); + final PyType parameterizedType = getParameterizedType(resolved, context); if (parameterizedType != null) { return parameterizedType; } @@ -328,22 +326,22 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { if (builtinCollection != null) { return builtinCollection; } - final PyType genericType = getGenericType(resolved, context, cache); + final PyType genericType = getGenericType(resolved, context); if (genericType != null) { return genericType; } - final Ref classType = getClassType(resolved, context); + final Ref classType = getClassType(resolved, context.getTypeContext()); if (classType != null) { return classType.get(); } - final PyType stringBasedType = getStringBasedType(resolved, context, cache); + final PyType stringBasedType = getStringBasedType(resolved, context); if (stringBasedType != null) { return stringBasedType; } return null; } finally { - cache.remove(resolved); + context.getExpressionCache().remove(resolved); } } @@ -371,18 +369,16 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { } @Nullable - public static PyType getTypeFromTargetExpression(@NotNull PyTargetExpression expression, - @NotNull TypeEvalContext context) { - return getTypeFromTargetExpression(expression, context, new HashSet<>()); + public static PyType getTypeFromTargetExpression(@NotNull PyTargetExpression expression, @NotNull TypeEvalContext context) { + return getTypeFromTargetExpression(expression, new Context(context)); } @Nullable private static PyType getTypeFromTargetExpression(@NotNull PyTargetExpression expression, - @NotNull TypeEvalContext context, - @NotNull Set cache) { + @NotNull Context context) { // XXX: Requires switching from stub to AST final PyExpression assignedValue = expression.findAssignedValue(); - return assignedValue != null ? getTypeForResolvedElement(assignedValue, context, cache) : null; + return assignedValue != null ? getTypeForResolvedElement(assignedValue, context) : null; } @Nullable @@ -407,17 +403,15 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { } @Nullable - private static Ref getOptionalType(@NotNull PsiElement element, - @NotNull TypeEvalContext context, - @NotNull Set cache) { + private static Ref getOptionalType(@NotNull PsiElement element, @NotNull Context context) { if (element instanceof PySubscriptionExpression) { final PySubscriptionExpression subscriptionExpr = (PySubscriptionExpression)element; final PyExpression operand = subscriptionExpr.getOperand(); - final Collection operandNames = resolveToQualifiedNames(operand, context); + final Collection operandNames = resolveToQualifiedNames(operand, context.getTypeContext()); if (operandNames.contains("typing.Optional")) { final PyExpression indexExpr = subscriptionExpr.getIndexExpression(); if (indexExpr != null) { - final PyType type = getType(indexExpr, context, cache); + final PyType type = getType(indexExpr, context); if (type != null) { return Ref.create(PyUnionType.union(type, PyNoneType.INSTANCE)); } @@ -429,20 +423,17 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { } @Nullable - private static PyType getStringBasedType(@NotNull PsiElement element, @NotNull TypeEvalContext context, @NotNull Set cache) { + private static PyType getStringBasedType(@NotNull PsiElement element, @NotNull Context context) { if (element instanceof PyStringLiteralExpression) { // XXX: Requires switching from stub to AST final String contents = ((PyStringLiteralExpression)element).getStringValue(); - return getStringBasedType(contents, element, context, cache); + return getStringBasedType(contents, element, context); } return null; } @Nullable - private static PyType getStringBasedType(@NotNull String contents, - @NotNull PsiElement anchor, - @NotNull TypeEvalContext context, - @NotNull Set cache) { + private static PyType getStringBasedType(@NotNull String contents, @NotNull PsiElement anchor, @NotNull Context context) { final Project project = anchor.getProject(); final PyExpressionCodeFragmentImpl codeFragment = new PyExpressionCodeFragmentImpl(project, "dummy.py", contents, false); codeFragment.setContext(anchor.getContainingFile()); @@ -453,21 +444,21 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { final PyTupleExpression tupleExpr = (PyTupleExpression)expr; final List elementTypes = new ArrayList(); for (PyExpression elementExpr : tupleExpr.getElements()) { - elementTypes.add(getType(elementExpr, context, cache)); + elementTypes.add(getType(elementExpr, context)); } return PyTupleType.create(anchor, elementTypes.toArray(new PyType[elementTypes.size()])); } - return getType(expr, context, cache); + return getType(expr, context); } return null; } @Nullable - private static PyType getCallableType(@NotNull PsiElement resolved, @NotNull TypeEvalContext context, @NotNull Set cache) { + private static PyType getCallableType(@NotNull PsiElement resolved, @NotNull Context context) { if (resolved instanceof PySubscriptionExpression) { final PySubscriptionExpression subscriptionExpr = (PySubscriptionExpression)resolved; final PyExpression operand = subscriptionExpr.getOperand(); - final Collection operandNames = resolveToQualifiedNames(operand, context); + final Collection operandNames = resolveToQualifiedNames(operand, context.getTypeContext()); if (operandNames.contains("typing.Callable")) { final PyExpression indexExpr = subscriptionExpr.getIndexExpression(); if (indexExpr instanceof PyTupleExpression) { @@ -479,10 +470,10 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { final List parameters = new ArrayList(); final PyListLiteralExpression listExpr = (PyListLiteralExpression)parametersExpr; for (PyExpression argExpr : listExpr.getElements()) { - parameters.add(new PyCallableParameterImpl(null, getType(argExpr, context, cache))); + parameters.add(new PyCallableParameterImpl(null, getType(argExpr, context))); } final PyExpression returnTypeExpr = elements[1]; - final PyType returnType = getType(returnTypeExpr, context, cache); + final PyType returnType = getType(returnTypeExpr, context); return new PyCallableTypeImpl(parameters, returnType); } } @@ -493,27 +484,25 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { } @Nullable - private static PyType getUnionType(@NotNull PsiElement element, @NotNull TypeEvalContext context, @NotNull Set cache) { + private static PyType getUnionType(@NotNull PsiElement element, @NotNull Context context) { if (element instanceof PySubscriptionExpression) { final PySubscriptionExpression subscriptionExpr = (PySubscriptionExpression)element; final PyExpression operand = subscriptionExpr.getOperand(); - final Collection operandNames = resolveToQualifiedNames(operand, context); + final Collection operandNames = resolveToQualifiedNames(operand, context.getTypeContext()); if (operandNames.contains("typing.Union")) { - return PyUnionType.union(getIndexTypes(subscriptionExpr, context, cache)); + return PyUnionType.union(getIndexTypes(subscriptionExpr, context)); } } return null; } @Nullable - private static PyGenericType getGenericType(@NotNull PsiElement element, - @NotNull TypeEvalContext context, - @NotNull Set cache) { + private static PyGenericType getGenericType(@NotNull PsiElement element, @NotNull Context context) { if (element instanceof PyCallExpression) { final PyCallExpression assignedCall = (PyCallExpression)element; final PyExpression callee = assignedCall.getCallee(); if (callee != null) { - final Collection calleeQNames = resolveToQualifiedNames(callee, context); + final Collection calleeQNames = resolveToQualifiedNames(callee, context.getTypeContext()); if (calleeQNames.contains("typing.TypeVar")) { final PyExpression[] arguments = assignedCall.getArguments(); if (arguments.length > 0) { @@ -521,7 +510,7 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { if (firstArgument instanceof PyStringLiteralExpression) { final String name = ((PyStringLiteralExpression)firstArgument).getStringValue(); if (name != null) { - return new PyGenericType(name, getGenericTypeBound(arguments, context, cache)); + return new PyGenericType(name, getGenericTypeBound(arguments, context)); } } } @@ -532,46 +521,40 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { } @Nullable - private static PyType getGenericTypeBound(@NotNull PyExpression[] typeVarArguments, - @NotNull TypeEvalContext context, - @NotNull Set cache) { - final List types = new ArrayList(); + private static PyType getGenericTypeBound(@NotNull PyExpression[] typeVarArguments, @NotNull Context context) { + final List types = new ArrayList<>(); for (int i = 1; i < typeVarArguments.length; i++) { - types.add(getType(typeVarArguments[i], context, cache)); + types.add(getType(typeVarArguments[i], context)); } return PyUnionType.union(types); } @NotNull - private static List getIndexTypes(@NotNull PySubscriptionExpression expression, - @NotNull TypeEvalContext context, - @NotNull Set cache) { + private static List getIndexTypes(@NotNull PySubscriptionExpression expression, @NotNull Context context) { final List types = new ArrayList(); final PyExpression indexExpr = expression.getIndexExpression(); if (indexExpr instanceof PyTupleExpression) { final PyTupleExpression tupleExpr = (PyTupleExpression)indexExpr; for (PyExpression expr : tupleExpr.getElements()) { - types.add(getType(expr, context, cache)); + types.add(getType(expr, context)); } } else if (indexExpr != null) { - types.add(getType(indexExpr, context, cache)); + types.add(getType(indexExpr, context)); } return types; } @Nullable - private static PyType getParameterizedType(@NotNull PsiElement element, - @NotNull TypeEvalContext context, - @NotNull Set cache) { + private static PyType getParameterizedType(@NotNull PsiElement element, @NotNull Context context) { if (element instanceof PySubscriptionExpression) { final PySubscriptionExpression subscriptionExpr = (PySubscriptionExpression)element; final PyExpression operand = subscriptionExpr.getOperand(); final PyExpression indexExpr = subscriptionExpr.getIndexExpression(); - final PyType operandType = getType(operand, context, cache); + final PyType operandType = getType(operand, context); if (operandType instanceof PyClassType) { final PyClass cls = ((PyClassType)operandType).getPyClass(); - final List indexTypes = getIndexTypes(subscriptionExpr, context, cache); + final List indexTypes = getIndexTypes(subscriptionExpr, context); if (PyNames.TUPLE.equals(cls.getQualifiedName())) { return PyTupleType.create(element, indexTypes.toArray(new PyType[indexTypes.size()])); } @@ -646,4 +629,23 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { } return null; } + + private static class Context { + @NotNull private final TypeEvalContext myContext; + @NotNull private final Set myCache = new HashSet<>(); + + private Context(@NotNull TypeEvalContext context) { + myContext = context; + } + + @NotNull + public TypeEvalContext getTypeContext() { + return myContext; + } + + @NotNull + public Set getExpressionCache() { + return myCache; + } + } }