PY-18386 Add recursion guard to handle recursive types via PEP 484 type aliases

Recursive type are not fully supported yet. It only prevents SOE for them.
This commit is contained in:
Mikhail Golubev
2016-03-09 17:24:52 +03:00
parent 9cb08375e6
commit 2c2cbf3adb
2 changed files with 133 additions and 66 deletions
@@ -25,6 +25,7 @@ import com.intellij.psi.PsiComment;
import com.intellij.psi.PsiElement;
import com.intellij.psi.PsiPolyVariantReference;
import com.intellij.psi.util.PsiTreeUtil;
import com.intellij.util.containers.HashSet;
import com.jetbrains.python.PyNames;
import com.jetbrains.python.psi.*;
import com.jetbrains.python.psi.impl.PyExpressionCodeFragmentImpl;
@@ -82,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);
final PyType type = getType(value, context, new HashSet<>());
if (type != null) {
final PyType optionalType = getOptionalTypeFromDefaultNone(param, type, context);
return Ref.create(optionalType != null ? optionalType : type);
@@ -126,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);
final PyType type = getType(value, context, new HashSet<>());
return type != null ? Ref.create(type) : null;
}
}
final PyType constructorType = getGenericConstructorType(function, context);
final PyType constructorType = getGenericConstructorType(function, context, new HashSet<>());
if (constructorType != null) {
return Ref.create(constructorType);
}
@@ -154,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);
return getType(typeExpr, context, new HashSet<>());
}
}
return null;
@@ -166,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);
final PyType type = getStringBasedType(comment, referenceTarget, context, new HashSet<>());
if (type instanceof PyTupleType) {
final PyTupleExpression tupleExpr = PsiTreeUtil.getParentOfType(target, PyTupleExpression.class);
if (tupleExpr != null) {
@@ -243,11 +244,11 @@ public class PyTypingTypeProvider extends PyTypeProviderBase {
}
@Nullable
private static PyType getGenericConstructorType(@NotNull PyFunction function, @NotNull TypeEvalContext context) {
private static PyType getGenericConstructorType(@NotNull PyFunction function, @NotNull TypeEvalContext context, @NotNull Set<PsiElement> cache) {
if (PyUtil.isInit(function)) {
final PyClass cls = function.getContainingClass();
if (cls != null) {
final List<PyGenericType> genericTypes = collectGenericTypes(cls, context);
final List<PyGenericType> genericTypes = collectGenericTypes(cls, context, cache);
final List<PyType> elementTypes = new ArrayList<PyType>(genericTypes);
if (!elementTypes.isEmpty()) {
return new PyCollectionTypeImpl(cls, false, elementTypes);
@@ -258,7 +259,7 @@ public class PyTypingTypeProvider extends PyTypeProviderBase {
}
@NotNull
private static List<PyGenericType> collectGenericTypes(@NotNull PyClass cls, @NotNull TypeEvalContext context) {
private static List<PyGenericType> collectGenericTypes(@NotNull PyClass cls, @NotNull TypeEvalContext context, @NotNull Set<PsiElement> cache) {
boolean isGeneric = false;
for (PyClass ancestor : cls.getAncestorClasses(context)) {
if (GENERIC_CLASSES.contains(ancestor.getQualifiedName())) {
@@ -274,7 +275,7 @@ public class PyTypingTypeProvider extends PyTypeProviderBase {
final PyExpression indexExpr = ((PySubscriptionExpression)expr).getIndexExpression();
if (indexExpr != null) {
for (PsiElement resolved : tryResolving(indexExpr, context)) {
final PyGenericType genericType = getGenericType(resolved, context);
final PyGenericType genericType = getGenericType(resolved, context, cache);
if (genericType != null) {
results.add(genericType);
}
@@ -288,49 +289,62 @@ public class PyTypingTypeProvider extends PyTypeProviderBase {
}
@Nullable
private static PyType getType(@NotNull PyExpression expression, @NotNull TypeEvalContext context) {
private static PyType getType(@NotNull PyExpression expression, @NotNull TypeEvalContext context, @NotNull Set<PsiElement> cache) {
final List<PyType> members = Lists.newArrayList();
for (PsiElement resolved : tryResolving(expression, context)) {
members.add(getTypeForResolvedElement(resolved, context));
members.add(getTypeForResolvedElement(resolved, context, cache));
}
return PyUnionType.union(members);
}
@Nullable
private static PyType getTypeForResolvedElement(@NotNull PsiElement resolved, @NotNull TypeEvalContext context) {
final PyType unionType = getUnionType(resolved, context);
if (unionType != null) {
return unionType;
private static PyType getTypeForResolvedElement(@NotNull PsiElement resolved,
@NotNull TypeEvalContext context,
@NotNull Set<PsiElement> cache) {
if (cache.contains(resolved)) {
// Recursive types are not yet supported
return null;
}
final Ref<PyType> optionalType = getOptionalType(resolved, context);
if (optionalType != null) {
return optionalType.get();
cache.add(resolved);
try {
final PyType unionType = getUnionType(resolved, context, cache);
if (unionType != null) {
return unionType;
}
final Ref<PyType> optionalType = getOptionalType(resolved, context, cache);
if (optionalType != null) {
return optionalType.get();
}
final PyType callableType = getCallableType(resolved, context, cache);
if (callableType != null) {
return callableType;
}
final PyType parameterizedType = getParameterizedType(resolved, context, cache);
if (parameterizedType != null) {
return parameterizedType;
}
final PyType builtinCollection = getBuiltinCollection(resolved);
if (builtinCollection != null) {
return builtinCollection;
}
final PyType genericType = getGenericType(resolved, context, cache);
if (genericType != null) {
return genericType;
}
final Ref<PyType> classType = getClassType(resolved, context);
if (classType != null) {
return classType.get();
}
final PyType stringBasedType = getStringBasedType(resolved, context, cache);
if (stringBasedType != null) {
return stringBasedType;
}
return null;
}
final PyType callableType = getCallableType(resolved, context);
if (callableType != null) {
return callableType;
finally {
cache.remove(resolved);
}
final PyType parameterizedType = getParameterizedType(resolved, context);
if (parameterizedType != null) {
return parameterizedType;
}
final PyType builtinCollection = getBuiltinCollection(resolved);
if (builtinCollection != null) {
return builtinCollection;
}
final PyType genericType = getGenericType(resolved, context);
if (genericType != null) {
return genericType;
}
final Ref<PyType> classType = getClassType(resolved, context);
if (classType != null) {
return classType.get();
}
final PyType stringBasedType = getStringBasedType(resolved, context);
if (stringBasedType != null) {
return stringBasedType;
}
return null;
}
@Nullable
@@ -357,10 +371,18 @@ public class PyTypingTypeProvider extends PyTypeProviderBase {
}
@Nullable
public static PyType getTypeFromTargetExpression(@NotNull PyTargetExpression expression, @NotNull TypeEvalContext context) {
public static PyType getTypeFromTargetExpression(@NotNull PyTargetExpression expression,
@NotNull TypeEvalContext context) {
return getTypeFromTargetExpression(expression, context, new HashSet<>());
}
@Nullable
private static PyType getTypeFromTargetExpression(@NotNull PyTargetExpression expression,
@NotNull TypeEvalContext context,
@NotNull Set<PsiElement> cache) {
// XXX: Requires switching from stub to AST
final PyExpression assignedValue = expression.findAssignedValue();
return assignedValue != null ? getTypeForResolvedElement(assignedValue, context) : null;
return assignedValue != null ? getTypeForResolvedElement(assignedValue, context, cache) : null;
}
@Nullable
@@ -385,7 +407,9 @@ public class PyTypingTypeProvider extends PyTypeProviderBase {
}
@Nullable
private static Ref<PyType> getOptionalType(@NotNull PsiElement element, @NotNull TypeEvalContext context) {
private static Ref<PyType> getOptionalType(@NotNull PsiElement element,
@NotNull TypeEvalContext context,
@NotNull Set<PsiElement> cache) {
if (element instanceof PySubscriptionExpression) {
final PySubscriptionExpression subscriptionExpr = (PySubscriptionExpression)element;
final PyExpression operand = subscriptionExpr.getOperand();
@@ -393,7 +417,7 @@ public class PyTypingTypeProvider extends PyTypeProviderBase {
if (operandNames.contains("typing.Optional")) {
final PyExpression indexExpr = subscriptionExpr.getIndexExpression();
if (indexExpr != null) {
final PyType type = getType(indexExpr, context);
final PyType type = getType(indexExpr, context, cache);
if (type != null) {
return Ref.create(PyUnionType.union(type, PyNoneType.INSTANCE));
}
@@ -405,17 +429,20 @@ public class PyTypingTypeProvider extends PyTypeProviderBase {
}
@Nullable
private static PyType getStringBasedType(@NotNull PsiElement element, @NotNull TypeEvalContext context) {
private static PyType getStringBasedType(@NotNull PsiElement element, @NotNull TypeEvalContext context, @NotNull Set<PsiElement> cache) {
if (element instanceof PyStringLiteralExpression) {
// XXX: Requires switching from stub to AST
final String contents = ((PyStringLiteralExpression)element).getStringValue();
return getStringBasedType(contents, element, context);
return getStringBasedType(contents, element, context, cache);
}
return null;
}
@Nullable
private static PyType getStringBasedType(@NotNull String contents, @NotNull PsiElement anchor, @NotNull TypeEvalContext context) {
private static PyType getStringBasedType(@NotNull String contents,
@NotNull PsiElement anchor,
@NotNull TypeEvalContext context,
@NotNull Set<PsiElement> cache) {
final Project project = anchor.getProject();
final PyExpressionCodeFragmentImpl codeFragment = new PyExpressionCodeFragmentImpl(project, "dummy.py", contents, false);
codeFragment.setContext(anchor.getContainingFile());
@@ -426,17 +453,17 @@ public class PyTypingTypeProvider extends PyTypeProviderBase {
final PyTupleExpression tupleExpr = (PyTupleExpression)expr;
final List<PyType> elementTypes = new ArrayList<PyType>();
for (PyExpression elementExpr : tupleExpr.getElements()) {
elementTypes.add(getType(elementExpr, context));
elementTypes.add(getType(elementExpr, context, cache));
}
return PyTupleType.create(anchor, elementTypes.toArray(new PyType[elementTypes.size()]));
}
return getType(expr, context);
return getType(expr, context, cache);
}
return null;
}
@Nullable
private static PyType getCallableType(@NotNull PsiElement resolved, @NotNull TypeEvalContext context) {
private static PyType getCallableType(@NotNull PsiElement resolved, @NotNull TypeEvalContext context, @NotNull Set<PsiElement> cache) {
if (resolved instanceof PySubscriptionExpression) {
final PySubscriptionExpression subscriptionExpr = (PySubscriptionExpression)resolved;
final PyExpression operand = subscriptionExpr.getOperand();
@@ -452,10 +479,10 @@ public class PyTypingTypeProvider extends PyTypeProviderBase {
final List<PyCallableParameter> parameters = new ArrayList<PyCallableParameter>();
final PyListLiteralExpression listExpr = (PyListLiteralExpression)parametersExpr;
for (PyExpression argExpr : listExpr.getElements()) {
parameters.add(new PyCallableParameterImpl(null, getType(argExpr, context)));
parameters.add(new PyCallableParameterImpl(null, getType(argExpr, context, cache)));
}
final PyExpression returnTypeExpr = elements[1];
final PyType returnType = getType(returnTypeExpr, context);
final PyType returnType = getType(returnTypeExpr, context, cache);
return new PyCallableTypeImpl(parameters, returnType);
}
}
@@ -466,20 +493,22 @@ public class PyTypingTypeProvider extends PyTypeProviderBase {
}
@Nullable
private static PyType getUnionType(@NotNull PsiElement element, @NotNull TypeEvalContext context) {
private static PyType getUnionType(@NotNull PsiElement element, @NotNull TypeEvalContext context, @NotNull Set<PsiElement> cache) {
if (element instanceof PySubscriptionExpression) {
final PySubscriptionExpression subscriptionExpr = (PySubscriptionExpression)element;
final PyExpression operand = subscriptionExpr.getOperand();
final Collection<String> operandNames = resolveToQualifiedNames(operand, context);
if (operandNames.contains("typing.Union")) {
return PyUnionType.union(getIndexTypes(subscriptionExpr, context));
return PyUnionType.union(getIndexTypes(subscriptionExpr, context, cache));
}
}
return null;
}
@Nullable
private static PyGenericType getGenericType(@NotNull PsiElement element, @NotNull TypeEvalContext context) {
private static PyGenericType getGenericType(@NotNull PsiElement element,
@NotNull TypeEvalContext context,
@NotNull Set<PsiElement> cache) {
if (element instanceof PyCallExpression) {
final PyCallExpression assignedCall = (PyCallExpression)element;
final PyExpression callee = assignedCall.getCallee();
@@ -492,7 +521,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));
return new PyGenericType(name, getGenericTypeBound(arguments, context, cache));
}
}
}
@@ -503,40 +532,46 @@ public class PyTypingTypeProvider extends PyTypeProviderBase {
}
@Nullable
private static PyType getGenericTypeBound(@NotNull PyExpression[] typeVarArguments, @NotNull TypeEvalContext context) {
private static PyType getGenericTypeBound(@NotNull PyExpression[] typeVarArguments,
@NotNull TypeEvalContext context,
@NotNull Set<PsiElement> cache) {
final List<PyType> types = new ArrayList<PyType>();
for (int i = 1; i < typeVarArguments.length; i++) {
types.add(getType(typeVarArguments[i], context));
types.add(getType(typeVarArguments[i], context, cache));
}
return PyUnionType.union(types);
}
@NotNull
private static List<PyType> getIndexTypes(@NotNull PySubscriptionExpression expression, @NotNull TypeEvalContext context) {
private static List<PyType> getIndexTypes(@NotNull PySubscriptionExpression expression,
@NotNull TypeEvalContext context,
@NotNull Set<PsiElement> cache) {
final List<PyType> types = new ArrayList<PyType>();
final PyExpression indexExpr = expression.getIndexExpression();
if (indexExpr instanceof PyTupleExpression) {
final PyTupleExpression tupleExpr = (PyTupleExpression)indexExpr;
for (PyExpression expr : tupleExpr.getElements()) {
types.add(getType(expr, context));
types.add(getType(expr, context, cache));
}
}
else if (indexExpr != null) {
types.add(getType(indexExpr, context));
types.add(getType(indexExpr, context, cache));
}
return types;
}
@Nullable
private static PyType getParameterizedType(@NotNull PsiElement element, @NotNull TypeEvalContext context) {
private static PyType getParameterizedType(@NotNull PsiElement element,
@NotNull TypeEvalContext context,
@NotNull Set<PsiElement> cache) {
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);
final PyType operandType = getType(operand, context, cache);
if (operandType instanceof PyClassType) {
final PyClass cls = ((PyClassType)operandType).getPyClass();
final List<PyType> indexTypes = getIndexTypes(subscriptionExpr, context);
final List<PyType> indexTypes = getIndexTypes(subscriptionExpr, context, cache);
if (PyNames.TUPLE.equals(cls.getQualifiedName())) {
return PyTupleType.create(element, indexTypes.toArray(new PyType[indexTypes.size()]));
}
@@ -472,6 +472,38 @@ public class PyTypingTest extends PyTestCase {
}
// PY-18386
public void testRecursiveType() {
doTest("Union[int, Any]",
"from typing import Union\n" +
"\n" +
"Type = Union[int, 'Type']\n" +
"expr = 42 # type: Type");
}
// PY-18386
public void testRecursiveType2() {
doTest("Dict[str, Union[Union[str, int, float], Any]]",
"from typing import Dict, Union\n" +
"\n" +
"JsonDict = Dict[str, Union[str, int, float, 'JsonDict']]\n" +
"\n" +
"def f(x: JsonDict):\n" +
" expr = x");
}
// PY-18386
public void testRecursiveType3() {
doTest("Union[Union[str, int], Any]",
"from typing import Union\n" +
"\n" +
"Type1 = Union[str, 'Type2']\n" +
"Type2 = Union[int, Type1]\n" +
"\n" +
"expr = None # type: Type1");
}
private void doTestNoInjectedText(@NotNull String text) {
myFixture.configureByText(PythonFileType.INSTANCE, text);
final InjectedLanguageManager languageManager = InjectedLanguageManager.getInstance(myFixture.getProject());