mirror of
https://gitflic.ru/project/openide/openide.git
synced 2026-09-27 10:03:11 +07:00
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:
@@ -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());
|
||||
|
||||
Reference in New Issue
Block a user