[python] Create PyParamSpecType together with other type parameters in PyTypingTypeProvider

GitOrigin-RevId: e80c9a06c47ad7fcbc80bb9b1b967da604e1b540
This commit is contained in:
Mikhail Golubev
2025-01-22 16:49:41 +00:00
committed by intellij-monorepo-bot
parent 36fe8a62a9
commit bfffa97843
@@ -908,10 +908,6 @@ public final class PyTypingTypeProvider extends PyTypeProviderWithCustomContext<
if (typeParameterType != null) {
return Ref.create(typeParameterType);
}
final PyType paramSpecType = getParamSpecType(resolved, context);
if (paramSpecType != null) {
return Ref.create(anchorTypeParameter(typeHint, paramSpecType, context));
}
final PyType callableParameterListType = getCallableParameterListType(resolved, context);
if (callableParameterListType != null) {
return Ref.create(callableParameterListType);
@@ -1478,22 +1474,24 @@ public final class PyTypingTypeProvider extends PyTypeProviderWithCustomContext<
final PyExpression callee = assignedCall.getCallee();
if (callee != null) {
final Collection<String> calleeQNames = resolveToQualifiedNames(callee, context.getTypeContext());
if (calleeQNames.contains(TYPE_VAR) || calleeQNames.contains(TYPE_VAR_EXT)
|| calleeQNames.contains(TYPE_VAR_TUPLE) || calleeQNames.contains(TYPE_VAR_TUPLE_EXT)) {
if (ContainerUtil.exists(calleeQNames, qname -> TYPE_PARAMETER_FACTORIES.contains(qname))) {
final PyExpression[] arguments = assignedCall.getArguments();
if (arguments.length > 0) {
final PyExpression firstArgument = arguments[0];
if (firstArgument instanceof PyStringLiteralExpression) {
final String name = ((PyStringLiteralExpression)firstArgument).getStringValue();
if (calleeQNames.contains(TYPE_VAR_TUPLE) || calleeQNames.contains(TYPE_VAR_TUPLE_EXT)) {
PyType defaultType = Ref.deref(getTypeVarDefaultType(assignedCall, context));
return new PyTypeVarTupleTypeImpl(name)
.withDefaultType(defaultType instanceof PyPositionalVariadicType positionalVariadicType
? Ref.create(positionalVariadicType) : null);
}
else {
return new PyTypeVarTypeImpl(name, getGenericTypeBound(arguments, context), getTypeVarDefaultType(assignedCall, context));
}
if (arguments.length > 0 && arguments[0] instanceof PyStringLiteralExpression nameArgument) {
final String name = nameArgument.getStringValue();
PyExpression defaultExpression = assignedCall.getKeywordArgument("default");
Ref<PyType> defaultType = defaultExpression != null ? getType(defaultExpression, context) : null;
if (calleeQNames.contains(TYPE_VAR_TUPLE) || calleeQNames.contains(TYPE_VAR_TUPLE_EXT)) {
return new PyTypeVarTupleTypeImpl(name)
.withDefaultType(defaultType != null && Ref.deref(defaultType) instanceof PyPositionalVariadicType posVariadic
? Ref.create(posVariadic) : null);
}
else if (calleeQNames.contains(TYPE_VAR) || calleeQNames.contains(TYPE_VAR_EXT)) {
return new PyTypeVarTypeImpl(name, getGenericTypeBound(arguments, context), defaultType);
}
else if (calleeQNames.contains(PARAM_SPEC) || calleeQNames.contains(PARAM_SPEC_EXT)) {
return new PyParamSpecType(name)
.withDefaultType(defaultType != null && Ref.deref(defaultType) instanceof PyCallableParameterVariadicType paramVariadic
? Ref.create(paramVariadic) : null);
}
}
}
@@ -1665,25 +1663,6 @@ public final class PyTypingTypeProvider extends PyTypeProviderWithCustomContext<
return Ref.create(Ref.deref(getType(starredExpression, context)));
}
private static @Nullable PyParamSpecType getParamSpecType(@NotNull PsiElement element, @NotNull Context context) {
if (!(element instanceof PyCallExpression assignedCall)) return null;
final var callee = assignedCall.getCallee();
if (callee == null) return null;
final var calleeQNames = resolveToQualifiedNames(callee, context.getTypeContext());
if (!calleeQNames.contains(PARAM_SPEC) && !calleeQNames.contains(PARAM_SPEC_EXT)) return null;
final var arguments = assignedCall.getArguments();
if (arguments.length == 0) return null;
final var firstArgument = arguments[0];
if (!(firstArgument instanceof PyStringLiteralExpression)) return null;
final var name = ((PyStringLiteralExpression)firstArgument).getStringValue();
return new PyParamSpecType(name).withDefaultType(getParamSpecDefaultType(assignedCall, context));
}
private static @Nullable PyType getGenericTypeBound(PyExpression @NotNull [] typeVarArguments, @NotNull Context context) {
final List<PyType> types = new ArrayList<>();
for (int i = 1; i < typeVarArguments.length; i++) {
@@ -1704,26 +1683,6 @@ public final class PyTypingTypeProvider extends PyTypeProviderWithCustomContext<
return PyUnionType.union(types);
}
private static @Nullable Ref<PyType> getTypeVarDefaultType(@NotNull PyCallExpression callExpression, @NotNull Context context) {
PyExpression defaultExpression = callExpression.getKeywordArgument("default");
if (defaultExpression != null) {
return getType(defaultExpression, context);
}
return null;
}
private static @Nullable Ref<PyCallableParameterVariadicType> getParamSpecDefaultType(@NotNull PyCallExpression callExpression,
@NotNull Context context) {
PyExpression defaultExpression = callExpression.getKeywordArgument("default");
if (defaultExpression != null) {
PyType defaultType = Ref.deref(getType(defaultExpression, context));
if (defaultType instanceof PyCallableParameterVariadicType variadicType) {
return Ref.create(variadicType);
}
}
return null;
}
private static @Nullable PyPositionalVariadicType getTypeVarTupleParameterDefaultType(@NotNull PyExpression defaultExpression, @NotNull Context context) {
return as(getTypeParameterBoundType(defaultExpression, context), PyPositionalVariadicType.class);
}