PY-71002 PEP-696: Implement support for defaults in Type Parameters

The changes include:
- Inferring implicitly parameterized types for classes if they are generic classes with type parameters that have defaults.
- Inferring partially parameterized types for classes where only some of the default Type Parameters are overridden by explicit parameterization.
- Supporting both old-style and new-style generic Type Aliases, now considering defaults of Type Parameters.

GitOrigin-RevId: c3a0d7eb4a85585df9638081291ff83b850eb7f6
This commit is contained in:
Daniil Kalinin
2024-09-07 11:11:13 +00:00
committed by intellij-monorepo-bot
parent 003cb6a891
commit cedd61e338
6 changed files with 319 additions and 29 deletions
@@ -726,6 +726,19 @@ public final class PyTypingTypeProvider extends PyTypeProviderWithCustomContext<
.toList();
}
@NotNull
private static List<PyTypeParameterType> collectTypeParametersFromTypeAliasStatement(@NotNull PyTypeAliasStatement typeAliasStatement, @NotNull Context context) {
PyTypeParameterList typeParameterList = typeAliasStatement.getTypeParameterList();
if (typeParameterList != null) {
List<PyTypeParameter> typeParameters = typeParameterList.getTypeParameters();
return StreamEx.of(typeParameters)
.map(typeParameter -> getTypeParameterTypeFromTypeParameter(typeParameter, context))
.nonNull()
.toList();
}
return Collections.emptyList();
}
public static boolean isGeneric(@NotNull PyWithAncestors descendant, @NotNull TypeEvalContext context) {
if (descendant instanceof PyClass pyClass && pyClass.getTypeParameterList() != null ||
descendant instanceof PyClassType pyClassType && pyClassType.getPyClass().getTypeParameterList() != null) {
@@ -835,6 +848,10 @@ public final class PyTypingTypeProvider extends PyTypeProviderWithCustomContext<
if (typeFromParenthesizedExpression != null) {
return typeFromParenthesizedExpression;
}
final PyType parameterizedTypeFromTypeAlias = getParameterizedTypeFromTypeAlias(alias, typeHint, resolved, context);
if (parameterizedTypeFromTypeAlias != null) {
return Ref.create(parameterizedTypeFromTypeAlias);
}
final PyType unionType = getUnionType(resolved, context);
if (unionType != null) {
return Ref.create(unionType);
@@ -1852,6 +1869,111 @@ public final class PyTypingTypeProvider extends PyTypeProviderWithCustomContext<
return types;
}
@Nullable
public static PyType tryParameterizeClassWithDefaults(@NotNull PyClassType classType,
@NotNull PyExpression anchor,
boolean toInstance,
@NotNull TypeEvalContext context) {
return staticWithCustomContext(context, customContext -> parameterizeClassWithDefaults(classType, anchor, toInstance, customContext));
}
@Nullable
private static PyType parameterizeClassWithDefaults(@NotNull PyClassType classType,
@NotNull PyExpression anchor,
boolean toInstance,
@NotNull Context context) {
PyClass pyClass = classType.getPyClass();
if (isGeneric(pyClass, context.getTypeContext())) {
PyCollectionType genericDefinitionType = PyTypeChecker.findGenericDefinitionType(pyClass, context.getTypeContext());
if (genericDefinitionType != null && ContainerUtil.exists(genericDefinitionType.getElementTypes(),
t -> t instanceof PyTypeParameterType typeParameterType &&
typeParameterType.getDefaultType() != null)) {
PyType parameterizedType;
if (anchor instanceof PySubscriptionExpression subscriptionExpression) {
List<PyType> indexTypes = getIndexTypes(subscriptionExpression, context);
parameterizedType = PyTypeChecker.parameterizeType(genericDefinitionType, indexTypes, context.myContext);
}
else {
List<PyTypeParameterType> typeParameters = collectTypeParameters(genericDefinitionType.getPyClass(), context);
parameterizedType = PyTypeChecker.trySubstituteByDefaultsOnly(genericDefinitionType, typeParameters, true, context.getTypeContext());
}
if (parameterizedType instanceof PyCollectionType collectionType) {
return toInstance ? collectionType.toInstance() : collectionType.toClass();
}
}
}
return null;
}
@Nullable
private static PyType getParameterizedTypeFromTypeAlias(@Nullable PyQualifiedNameOwner alias,
@NotNull PsiElement typeHint,
@NotNull PsiElement element,
@NotNull Context context) {
if (alias instanceof PyTypeAliasStatement typeAliasStatement) {
return getParameterizedTypeFromTypeAliasStatement(typeAliasStatement, typeHint, element, context);
}
if (alias != null && element instanceof PyExpression assignedExpression) {
if (typeHint instanceof PySubscriptionExpression subscriptionExpr) {
List<PyType> indexTypes = getIndexTypes(subscriptionExpr, context);
PyType assignedType = Ref.deref(getType(assignedExpression, context));
if (assignedType != null) {
return PyTypeChecker.parameterizeType(assignedType, indexTypes, context.myContext);
}
}
if (typeHint instanceof PyReferenceExpression) {
PyType expressionType = Ref.deref(getType(assignedExpression, context));
if (expressionType != null && !(expressionType instanceof PyTypeParameterType)) {
List<PyTypeParameterType> typeAliasTypeParams =
PyTypeChecker.collectGenerics(expressionType, context.getTypeContext()).getAllTypeParameters();
if (!typeAliasTypeParams.isEmpty()) {
return PyTypeChecker.trySubstituteByDefaultsOnly(expressionType, typeAliasTypeParams, false, context.getTypeContext());
}
}
if (expressionType instanceof PyCollectionType) {
return expressionType;
}
}
}
return null;
}
@Nullable
private static PyType getParameterizedTypeFromTypeAliasStatement(@NotNull PyTypeAliasStatement typeAliasStatement,
@NotNull PsiElement typeHint,
@NotNull PsiElement element,
@NotNull Context context) {
if (element instanceof PyExpression expression) {
if (typeHint instanceof PySubscriptionExpression subscriptionExpression) {
final List<PyType> indexTypes = getIndexTypes(subscriptionExpression, context);
PyType assignedType = Ref.deref(getType(expression, context));
List<PyTypeParameterType> typeParamsFromAliasStmt = collectTypeParametersFromTypeAliasStatement(typeAliasStatement, context);
if (assignedType != null && !typeParamsFromAliasStmt.isEmpty()) {
PyTypeChecker.GenericSubstitutions substitutions =
PyTypeChecker.mapTypeParametersToSubstitutions(new PyTypeChecker.GenericSubstitutions(),
typeParamsFromAliasStmt,
indexTypes,
Option.MAP_UNMATCHED_EXPECTED_TYPES_TO_ANY);
if (substitutions == null) return null;
var substitutionsWithDefaults = PyTypeChecker.getSubstitutionsWithDefaults(substitutions);
return PyTypeChecker.substitute(assignedType, substitutionsWithDefaults, context.getTypeContext());
}
}
else if (typeHint instanceof PyReferenceExpression) {
List<PyTypeParameterType> typeAliasTypeParams = collectTypeParametersFromTypeAliasStatement(typeAliasStatement, context);
if (!typeAliasTypeParams.isEmpty()) {
PyType operandType = Ref.deref(getType(expression, context));
if (operandType != null) {
return PyTypeChecker.trySubstituteByDefaultsOnly(operandType, typeAliasTypeParams, false, context.getTypeContext());
}
}
}
}
return null;
}
@Nullable
private static PyType getParameterizedType(@NotNull PsiElement element, @NotNull Context context) {
if (element instanceof PySubscriptionExpression subscriptionExpr) {
@@ -1953,6 +2075,16 @@ public final class PyTypingTypeProvider extends PyTypeProviderWithCustomContext<
}
}
}
if (expression instanceof PySubscriptionExpression subscriptionExpr) {
// Possibly a parameterized type alias
PyExpression operandExpression = subscriptionExpr.getOperand();
List<Pair<PyQualifiedNameOwner, PsiElement>> results = tryResolvingWithAliases(operandExpression, context);
for (Pair<PyQualifiedNameOwner, PsiElement> pair : results) {
if (pair.getFirst() != null && pair.getSecond() != null) {
elements.add(Pair.create(pair.getFirst(), pair.getSecond()));
}
}
}
return !elements.isEmpty() ? elements : Collections.singletonList(Pair.create(null, expression));
}
@@ -539,7 +539,15 @@ public final class PyCallExpressionHelper {
}
}
}
if (callee instanceof PySubscriptionExpression) {
if (callee instanceof PySubscriptionExpression subscriptionExpression) {
PyExpression operandExpr = subscriptionExpression.getOperand();
PyType operandType = context.getType(operandExpr);
if (operandType instanceof PyClassType classType) {
PyType parameterizeClassWithDefaults = PyTypingTypeProvider.tryParameterizeClassWithDefaults(classType, callee, true, context);
if (parameterizeClassWithDefaults instanceof PyCollectionType) {
return parameterizeClassWithDefaults;
}
}
final PyType parametrizedType = Ref.deref(PyTypingTypeProvider.getType(callee, context));
if (parametrizedType != null) {
return parametrizedType;
@@ -650,9 +658,19 @@ public final class PyCallExpressionHelper {
return new PyCollectionTypeImpl(receiverClass, false, elementTypes);
}
if (initOrNewCallType instanceof PyClassType classType) {
PyType implicitlyParameterized = PyTypingTypeProvider.tryParameterizeClassWithDefaults(classType, callSite, true, context);
if (implicitlyParameterized instanceof PyCollectionType collectionType) {
return collectionType.toInstance();
}
}
return new PyClassTypeImpl(receiverClass, false);
}
if (initOrNewCallType instanceof PyCollectionType) {
return initOrNewCallType;
}
if (initOrNewCallType == null) {
return PyUnionType.createWeakType(new PyClassTypeImpl(receiverClass, false));
}
@@ -223,7 +223,8 @@ public class PyFunctionImpl extends PyBaseElementImpl<PyFunctionStub> implements
}
final var substitutionsWithUnresolvedReturnGenerics =
PyTypeChecker.getSubstitutionsWithUnresolvedReturnGenerics(getParameters(context), type, substitutions, context);
type = PyTypeChecker.substitute(type, substitutionsWithUnresolvedReturnGenerics, context);
final var substitutionsWithDefaults = PyTypeChecker.getSubstitutionsWithDefaults(substitutionsWithUnresolvedReturnGenerics);
type = PyTypeChecker.substitute(type, substitutionsWithDefaults, context);
}
else {
type = null;
@@ -18,6 +18,7 @@ import com.jetbrains.python.codeInsight.controlflow.PyTypeAssertionEvaluator;
import com.jetbrains.python.codeInsight.controlflow.ReadWriteInstruction;
import com.jetbrains.python.codeInsight.controlflow.ScopeOwner;
import com.jetbrains.python.codeInsight.dataflow.scope.ScopeUtil;
import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider;
import com.jetbrains.python.psi.*;
import com.jetbrains.python.psi.impl.references.PyImportReference;
import com.jetbrains.python.psi.impl.references.PyQualifiedReference;
@@ -359,10 +360,21 @@ public class PyReferenceExpressionImpl extends PyElementImpl implements PyRefere
PyType qualifierType = context.getType(qualifier);
boolean possiblyParameterizedQualifier = !(qualifierType instanceof PyModuleType || qualifierType instanceof PyImportedModuleType);
if (possiblyParameterizedQualifier && PyTypeChecker.hasGenerics(type, context)) {
final var substitutions =
PyTypeChecker.unifyGenericCall(qualifier, Collections.emptyMap(), context);
if (qualifierType instanceof PyCollectionType collectionType && collectionType.isDefinition()) {
PyCollectionType genericDefinitionType = PyTypeChecker.findGenericDefinitionType(collectionType.getPyClass(), context);
if (genericDefinitionType != null && type != null) {
List<PyTypeParameterType> typeParameterTypes =
ContainerUtil.filterIsInstance(genericDefinitionType.getElementTypes(), PyTypeParameterType.class);
PyType typeWithSubstitutions = PyTypeChecker.trySubstituteByDefaultsOnly(type, typeParameterTypes, true, context);
if (typeWithSubstitutions != null) {
return typeWithSubstitutions;
}
}
}
final var substitutions = PyTypeChecker.unifyGenericCall(qualifier, Collections.emptyMap(), context);
if (substitutions != null) {
final PyType substituted = PyTypeChecker.substitute(type, substitutions, context);
var substitutionsWithDefaults = PyTypeChecker.getSubstitutionsWithDefaults(substitutions);
final PyType substituted = PyTypeChecker.substitute(type, substitutionsWithDefaults, context);
if (substituted != null) {
return substituted;
}
@@ -371,6 +383,13 @@ public class PyReferenceExpressionImpl extends PyElementImpl implements PyRefere
}
}
if (type instanceof PyClassType classType && !(type instanceof PyCollectionType)) {
PyType parameterizedType = PyTypingTypeProvider.tryParameterizeClassWithDefaults(classType, anchor, false, context);
if (parameterizedType instanceof PyCollectionType collectionType) {
return collectionType;
}
}
return type;
}
@@ -18,15 +18,13 @@ package com.jetbrains.python.psi.impl;
import com.intellij.lang.ASTNode;
import com.intellij.psi.PsiPolyVariantReference;
import com.intellij.psi.PsiReference;
import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider;
import com.jetbrains.python.psi.PyElementVisitor;
import com.jetbrains.python.psi.PyExpression;
import com.jetbrains.python.psi.PySubscriptionExpression;
import com.jetbrains.python.psi.impl.references.PyOperatorReference;
import com.jetbrains.python.psi.resolve.PyResolveContext;
import com.jetbrains.python.psi.types.PyTupleType;
import com.jetbrains.python.psi.types.PyType;
import com.jetbrains.python.psi.types.PyTypedDictType;
import com.jetbrains.python.psi.types.TypeEvalContext;
import com.jetbrains.python.psi.types.*;
import org.jetbrains.annotations.NotNull;
import org.jetbrains.annotations.Nullable;
@@ -61,6 +59,12 @@ public class PySubscriptionExpressionImpl extends PyElementImpl implements PySub
.map(typedDictType::getElementType)
.orElse(null);
}
if (type instanceof PyClassType classType) {
PyType parameterizedType = PyTypingTypeProvider.tryParameterizeClassWithDefaults(classType, this, false, context);
if (parameterizedType instanceof PyCollectionType collectionType) {
return collectionType.toClass();
}
}
return PyCallExpressionHelper.getCallType(this, context, key);
}
@@ -961,10 +961,17 @@ public final class PyTypeChecker {
boolean canGetBoundFromArguments = typeParamsFromParameterTypes.paramSpecs.contains(paramSpecType);
boolean isAlreadyBound = existingSubstitutions.paramSpecs.containsKey(paramSpecType);
if (canGetBoundFromArguments && !isAlreadyBound) {
existingSubstitutions.paramSpecs.put(paramSpecType, new PyParamSpecType(paramSpecType.getName())
.withParameters(List.of(PyCallableParameterImpl.positionalNonPsi("args", null),
PyCallableParameterImpl.keywordNonPsi("kwargs", null)),
context));
if (paramSpecType.getDefaultType() != null) {
PyType defaultType = paramSpecType.getDefaultType();
if (defaultType instanceof PyParamSpecType defaultParamSpec) {
existingSubstitutions.paramSpecs.put(paramSpecType, defaultParamSpec);
}
} else {
existingSubstitutions.paramSpecs.put(paramSpecType, new PyParamSpecType(paramSpecType.getName())
.withParameters(List.of(PyCallableParameterImpl.positionalNonPsi("args", null),
PyCallableParameterImpl.keywordNonPsi("kwargs", null)),
context));
}
}
}
return existingSubstitutions;
@@ -1042,6 +1049,12 @@ public final class PyTypeChecker {
}
collectGenerics(callable.getReturnType(context), context, generics, visited);
}
else if (type instanceof PyConcatenateType concatenateType) {
for (PyType elementType : concatenateType.getFirstTypes()) {
collectGenerics(elementType, context, generics, visited);
}
collectGenerics(concatenateType.getParamSpec(), context, generics, visited);
}
else if (type instanceof PyUnpackedTupleType unpackedTupleType) {
for (PyType elementType : unpackedTupleType.getElementTypes()) {
collectGenerics(elementType, context, generics, visited);
@@ -1108,6 +1121,17 @@ public final class PyTypeChecker {
//
// where we want to prevent replacing T with Any, even though it cannot be inferred at a call site.
if (!substitutions.typeVars.containsKey(typeVar) && !substitutions.typeVars.containsKey(invert(typeVar))) {
PyTypeVarType substitution = StreamEx.of(substitutions.typeVars.keySet())
.findFirst(typeVarType -> {
return typeVarType.getDeclarationElement() != null
&& (typeVar.getScopeOwner() == null || (typeVarType.getScopeOwner() == typeVar.getScopeOwner()))
&& typeVarType.getDeclarationElement().equals(typeVar.getDeclarationElement());
})
.orElse(null);
if (substitution != null) {
return substitute(substitution, substitutions, context, substituting);
}
return typeVar;
}
PyType substitution = substitutions.typeVars.get(typeVar);
@@ -1118,6 +1142,26 @@ public final class PyTypeChecker {
substitution = invert(invertedSubstitution);
}
}
if (substitution instanceof PyTypeVarType) {
PyQualifiedNameOwner scope = typeVar.getScopeOwner();
if (scope instanceof PyClass pyClass) {
PyType genericTypeFromClass = getGenericTypeForClass(context, pyClass);
if (genericTypeFromClass instanceof PyCollectionType) {
Generics generics = collectGenerics(genericTypeFromClass, context);
PyType finalSubstitution = substitution;
PyTypeVarType sameScopeSubstitution = StreamEx.of(generics.typeVars)
.findFirst(typeVarType -> {
return typeVarType.getDeclarationElement() != null
&& typeVarType.getDeclarationElement().equals(finalSubstitution.getDeclarationElement());
}).orElse(null);
if (sameScopeSubstitution != null && sameScopeSubstitution.getDefaultType() != null) {
return substitute(sameScopeSubstitution, substitutions, context, substituting);
}
}
}
}
// TODO remove !typeVar.equals(substitution) part, it's necessary due to the logic in unifyReceiverWithParamSpecs
if (!typeVar.equals(substitution) && hasGenerics(substitution, context)) {
return substitute(substitution, substitutions, context, substituting);
@@ -1347,6 +1391,41 @@ public final class PyTypeChecker {
return match(expectedType, actualType, context, substitutions);
}
@NotNull
public static GenericSubstitutions getSubstitutionsWithDefaults(@NotNull GenericSubstitutions substitutions) {
Map<PyTypeVarType, PyType> typeVersWithDefaults = new HashMap<>();
Map<PyParamSpecType, PyParamSpecType> paramSpecsWithDefaults = new HashMap<>();
Map<PyTypeVarTupleType, PyVariadicType> typeVarTuplesWithDefaults = new HashMap<>();
substitutions.getTypeVars().forEach((key, value) -> {
if (key.getDefaultType() != null && value == null) {
typeVersWithDefaults.put(key, key.getDefaultType());
}
else if (value instanceof PyTypeVarType typeVarType
&& typeVarType.getDefaultType() != null
&& !typeVarType.getDefaultType().equals(key)) {
typeVersWithDefaults.put(key, typeVarType.getDefaultType());
}
});
substitutions.getParamSpecs().forEach((key, value) -> {
if (key.getDefaultType() instanceof PyParamSpecType defaultType
&& (value == null || key.equals(value))) {
paramSpecsWithDefaults.put(key, defaultType);
}
});
substitutions.getTypeVarTuples().forEach((key, value) -> {
if (key.getDefaultType() instanceof PyVariadicType variadicType
&& (value == null
|| (value instanceof PyUnpackedTupleType unpackedTupleType && unpackedTupleType.getElementTypes().isEmpty())
|| key.equals(value))) {
typeVarTuplesWithDefaults.put(key, variadicType);
}
});
substitutions.typeVars.putAll(typeVersWithDefaults);
substitutions.paramSpecs.putAll(paramSpecsWithDefaults);
substitutions.typeVarTuples.putAll(typeVarTuplesWithDefaults);
return substitutions;
}
private static boolean matchContainer(@Nullable PyCallableParameter container, @NotNull List<? extends PyExpression> arguments,
@NotNull GenericSubstitutions substitutions, @NotNull TypeEvalContext context) {
if (container == null) {
@@ -1557,21 +1636,13 @@ public final class PyTypeChecker {
@NotNull TypeEvalContext context) {
Generics typeParams = collectGenerics(genericType, context);
if (!typeParams.isEmpty()) {
var substitutions = new GenericSubstitutions();
List<PyType> expectedTypeParams = new ArrayList<>(new LinkedHashSet<>(typeParams.getAllTypeParameters()));
PyTypeParameterMapping mapping = PyTypeParameterMapping.mapByShape(expectedTypeParams, actualTypeParams, Option.MAP_UNMATCHED_EXPECTED_TYPES_TO_ANY);
if (mapping != null) {
for (Couple<PyType> pair : mapping.getMappedTypes()) {
if (pair.getFirst() instanceof PyTypeVarType typeVar) {
substitutions.typeVars.put(typeVar, pair.getSecond());
}
else if (pair.getFirst() instanceof PyTypeVarTupleType typeVarTuple) {
assert pair.getSecond() instanceof PyVariadicType;
substitutions.typeVarTuples.put(typeVarTuple, (PyVariadicType)pair.getSecond());
}
}
}
return substitute(genericType, substitutions, context);
var substitutions = mapTypeParametersToSubstitutions(new GenericSubstitutions(),
expectedTypeParams, actualTypeParams,
Option.MAP_UNMATCHED_EXPECTED_TYPES_TO_ANY);
if (substitutions == null) return null;
var substitutionsWithDefaults = getSubstitutionsWithDefaults(substitutions);
return substitute(genericType, substitutionsWithDefaults, context);
}
// An already parameterized type, don't override existing values for type parameters
else if (genericType instanceof PyCollectionType) {
@@ -1581,9 +1652,54 @@ public final class PyTypeChecker {
PyClass cls = ((PyClassType)genericType).getPyClass();
return new PyCollectionTypeImpl(cls, false, actualTypeParams);
}
else {
return null;
return null;
}
@ApiStatus.Internal
@Nullable
public static GenericSubstitutions mapTypeParametersToSubstitutions(@NotNull GenericSubstitutions substitutions,
@NotNull List<? extends PyType> expectedTypes,
@NotNull List<? extends PyType> actualTypes,
Option @NotNull ... options) {
PyTypeParameterMapping mapping = PyTypeParameterMapping.mapByShape(expectedTypes, actualTypes, options);
if (mapping != null) {
return fillSubstitutionsWithTypeParameters(substitutions, mapping.getMappedTypes());
}
return null;
}
@NotNull
@ApiStatus.Internal
public static GenericSubstitutions fillSubstitutionsWithTypeParameters(@NotNull GenericSubstitutions substitutions, @NotNull List<Couple<PyType>> typeParameters) {
for (Couple<PyType> pair : typeParameters) {
if (pair.getFirst() instanceof PyTypeVarType typeVar) {
substitutions.typeVars.put(typeVar, pair.getSecond());
}
else if (pair.getFirst() instanceof PyTypeVarTupleType typeVarTuple) {
substitutions.typeVarTuples.put(typeVarTuple, as(pair.getSecond(), PyVariadicType.class));
}
else if (pair.getFirst() instanceof PyParamSpecType paramSpec) {
substitutions.paramSpecs.put(paramSpec, as(pair.getSecond(), PyParamSpecType.class));
}
}
return substitutions;
}
@ApiStatus.Internal
@Nullable
public static PyType trySubstituteByDefaultsOnly(@NotNull PyType targetType,
List<PyTypeParameterType> typeParameterTypes,
boolean unmappedToAny,
@NotNull TypeEvalContext context) {
if (!typeParameterTypes.isEmpty()) {
List<Couple<PyType>> typeParametersToSubstitutions = ContainerUtil.map(typeParameterTypes, typeParameterType
-> new Couple<>(typeParameterType, unmappedToAny ? null : typeParameterType));
var substitutions = fillSubstitutionsWithTypeParameters(new GenericSubstitutions(), typeParametersToSubstitutions);
var subsWithDefaults = getSubstitutionsWithDefaults(substitutions);
return substitute(targetType, subsWithDefaults, context);
}
return null;
}
@ApiStatus.Internal