PY-77538 Introduce PyCallableParameterListType and PyCallableParameterVariadicType interface

PyCallableParameterListType type represents a concrete list of PyCallableParameters that
a PyParamSpecType can be specialized with.

Now PyParamSpecType and PyCallableParameterListType are similar to PyTypeVarTupleType and
PyUnpackedTupleType (a type parameter and its concrete specialization).

PyCallableParameterVariadicType indicates types that PyParamSpecType can be specialized with.
Namely, another PyParamSpecType, PyCallableParameterListType and PyConcatenateType.

GitOrigin-RevId: cc254b64884e637c1200f6334b6680ea3444bb8a
This commit is contained in:
Mikhail Golubev
2024-11-21 11:26:23 +00:00
committed by intellij-monorepo-bot
parent 816845add5
commit b0bb5138f7
13 changed files with 164 additions and 169 deletions
@@ -0,0 +1,16 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.jetbrains.python.psi.types;
import org.jetbrains.annotations.ApiStatus;
import org.jetbrains.annotations.NotNull;
import java.util.List;
/**
* Represents a series of {@link PyCallableParameter} used either as a part of {@link PyCallableType} or a substitution
* for {@link PyParamSpecType}.
*/
@ApiStatus.Experimental
public interface PyCallableParameterListType extends PyCallableParameterVariadicType {
@NotNull List<PyCallableParameter> getParameters();
}
@@ -0,0 +1,11 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.jetbrains.python.psi.types;
import org.jetbrains.annotations.ApiStatus;
/**
* Represents an entity in the type system that stands for a parameter list of a callable type.
*/
@ApiStatus.Experimental
public non-sealed interface PyCallableParameterVariadicType extends PyVariadicType {
}
@@ -5,6 +5,9 @@ import org.jetbrains.annotations.ApiStatus;
/**
* A type representing an ordered series of other types.
* It can appear only as a type parameter/argument of another generic type.
* <p>
* Two variants of such types described in <a href="https://peps.python.org/pep-0646/">PEP 646 – Variadic Generics</a> are
* TypeVarTuples and unpacked tuple types.
*
@@ -12,6 +15,6 @@ import org.jetbrains.annotations.ApiStatus;
* @see PyUnpackedTupleType
*/
@ApiStatus.Experimental
public interface PyPositionalVariadicType extends PyVariadicType {
public sealed interface PyPositionalVariadicType extends PyVariadicType permits PyTypeVarTupleType, PyUnpackedTupleType {
}
@@ -23,5 +23,5 @@ package com.jetbrains.python.psi.types;
* @see <a href="https://peps.python.org/pep-0646/#type-variable-tuples">PEP 646 – Variadic Generics</a>
* @see PyUnpackedTupleType
*/
public interface PyTypeVarTupleType extends PyTypeParameterType, PyPositionalVariadicType {
public non-sealed interface PyTypeVarTupleType extends PyTypeParameterType, PyPositionalVariadicType {
}
@@ -17,7 +17,7 @@ import java.util.List;
* @see <a href="https://peps.python.org/pep-0646/#unpacking-tuple-types">PEP 646 – Variadic Generics</a>
* @see PyTypeVarTupleType
*/
public interface PyUnpackedTupleType extends PyPositionalVariadicType {
public non-sealed interface PyUnpackedTupleType extends PyPositionalVariadicType {
/**
* Returns types contained inside this unpacked tuple type.
* <p>
@@ -22,7 +22,7 @@ import java.util.List;
* @see PyPositionalVariadicType
*/
@ApiStatus.Experimental
public interface PyVariadicType extends PyType {
public sealed interface PyVariadicType extends PyType permits PyPositionalVariadicType, PyCallableParameterVariadicType {
@Override
default boolean isBuiltin() {
return false;
@@ -1649,7 +1649,8 @@ public final class PyTypingTypeProvider extends PyTypeProviderWithCustomContext<
.withDeclarationElement(declarationElement);
}
case ParamSpec -> {
PyType defaultType = defaultExpression != null ? createParamSpecDefaultTypeFromExpression(name, defaultExpression, context) : null;
PyCallableParameterVariadicType defaultType =
defaultExpression != null ? createParamSpecDefaultTypeFromExpression(defaultExpression, context) : null;
yield new PyParamSpecType(name)
.withScopeOwner(scopeOwner)
.withDefaultType(defaultType)
@@ -1796,7 +1797,7 @@ public final class PyTypingTypeProvider extends PyTypeProviderWithCustomContext<
if (!(firstArgument instanceof PyStringLiteralExpression)) return null;
final var name = ((PyStringLiteralExpression)firstArgument).getStringValue();
return new PyParamSpecType(name).withDefaultType(getParamSpecDefaultType(name, assignedCall, context));
return new PyParamSpecType(name).withDefaultType(getParamSpecDefaultType(assignedCall, context));
}
@Nullable
@@ -1830,18 +1831,18 @@ public final class PyTypingTypeProvider extends PyTypeProviderWithCustomContext<
}
@Nullable
private static PyType getParamSpecDefaultType(@NotNull String name, @NotNull PyCallExpression callExpression, @NotNull Context context) {
private static PyCallableParameterVariadicType getParamSpecDefaultType(@NotNull PyCallExpression callExpression,
@NotNull Context context) {
PyExpression defaultExpression = callExpression.getKeywordArgument("default");
if (defaultExpression != null) {
return createParamSpecDefaultTypeFromExpression(name, defaultExpression, context);
return createParamSpecDefaultTypeFromExpression(defaultExpression, context);
}
return null;
}
@Nullable
private static PyParamSpecType createParamSpecDefaultTypeFromExpression(@NotNull String name,
@NotNull PyExpression expression,
@NotNull Context context) {
private static PyCallableParameterVariadicType createParamSpecDefaultTypeFromExpression(@NotNull PyExpression expression,
@NotNull Context context) {
if (expression instanceof PyListLiteralExpression listLiteralExpression) {
PyExpression[] defaultExpressions = listLiteralExpression.getElements();
List<PyType> defaultArgumentTypes = StreamEx.of(defaultExpressions)
@@ -1849,9 +1850,7 @@ public final class PyTypingTypeProvider extends PyTypeProviderWithCustomContext<
.map(defExpr -> Ref.deref(getType(defExpr, context)))
.toList();
return new PyParamSpecType(name)
.withParameters(ContainerUtil
.map(defaultArgumentTypes, argType -> PyCallableParameterImpl.nonPsi(argType)), context.getTypeContext());
return new PyCallableParameterListTypeImpl(ContainerUtil.map(defaultArgumentTypes, PyCallableParameterImpl::nonPsi));
}
if (expression instanceof PyReferenceExpression) {
PyType referenceType = Ref.deref(getType(expression, context));
@@ -1885,7 +1884,7 @@ public final class PyTypingTypeProvider extends PyTypeProviderWithCustomContext<
if (indexExpr instanceof PyTupleExpression tupleExpr) {
for (PyExpression expr : tupleExpr.getElements()) {
if (expr instanceof PyListLiteralExpression listLiteralExpression) { // ParamSpec explicit parameterization
types.add(createParamSpecDefaultTypeFromExpression("P", listLiteralExpression, context));
types.add(createParamSpecDefaultTypeFromExpression(listLiteralExpression, context));
}
else {
types.add(Ref.deref(getType(expr, context)));
@@ -1893,7 +1892,7 @@ public final class PyTypingTypeProvider extends PyTypeProviderWithCustomContext<
}
}
else if (indexExpr instanceof PyListLiteralExpression listLiteralExpression) {
types.add(createParamSpecDefaultTypeFromExpression("P", listLiteralExpression, context));
types.add(createParamSpecDefaultTypeFromExpression(listLiteralExpression, context));
}
else if (indexExpr != null) {
types.add(Ref.deref(getType(indexExpr, context)));
@@ -407,11 +407,10 @@ public class PyTypeCheckerInspection extends PyInspection {
@NotNull List<AnalyzeArgumentResult> result,
@NotNull List<UnexpectedArgumentForParamSpec> unexpectedArgumentForParamSpecs,
@NotNull List<UnfilledParameterFromParamSpec> unfilledParameterFromParamSpecs) {
paramSpec = Objects.requireNonNullElse(substitutions.getParamSpecs().get(paramSpec), paramSpec);
List<PyCallableParameter> parameters = paramSpec.getParameters();
if (parameters == null) return;
PyCallableParameterListType paramSpecSubst = as(substitutions.getParamSpecs().get(paramSpec), PyCallableParameterListType.class);
if (paramSpecSubst == null) return;
var mapping = analyzeArguments(arguments, parameters, myTypeEvalContext);
var mapping = analyzeArguments(arguments, paramSpecSubst.getParameters(), myTypeEvalContext);
for (var item: mapping.getMappedParameters().entrySet()) {
PyExpression argument = item.getKey();
PyCallableParameter parameter = item.getValue();
@@ -0,0 +1,45 @@
package com.jetbrains.python.psi.types;
import com.intellij.openapi.util.text.StringUtil;
import com.jetbrains.python.PyNames;
import org.jetbrains.annotations.NotNull;
import java.util.List;
import java.util.Objects;
public final class PyCallableParameterListTypeImpl implements PyCallableParameterListType {
private final List<PyCallableParameter> myParameters;
public PyCallableParameterListTypeImpl(@NotNull List<PyCallableParameter> parameters) {
myParameters = List.copyOf(parameters);
}
@Override
public @NotNull List<PyCallableParameter> getParameters() {
return myParameters;
}
@Override
public @NotNull String getName() {
final TypeEvalContext context = TypeEvalContext.codeInsightFallback(null);
return String.format("[%s]",
StringUtil.join(myParameters, param -> {
PyType type = param.getType(context);
return type != null ? type.getName() : PyNames.UNKNOWN_TYPE;
},
", "));
}
@Override
public boolean equals(Object o) {
if (this == o) return true;
if (o == null || getClass() != o.getClass()) return false;
PyCallableParameterListTypeImpl type = (PyCallableParameterListTypeImpl)o;
return Objects.equals(myParameters, type.myParameters);
}
@Override
public int hashCode() {
return Objects.hash(myParameters);
}
}
@@ -1,31 +1,8 @@
package com.jetbrains.python.psi.types
import com.intellij.psi.PsiElement
import com.intellij.util.ProcessingContext
import com.jetbrains.python.psi.AccessDirection
import com.jetbrains.python.psi.PyExpression
import com.jetbrains.python.psi.resolve.PyResolveContext
import com.jetbrains.python.psi.resolve.RatedResolveResult
/**
* Type of typing.Concatenate to store corresponding first type and parameter specification
*/
class PyConcatenateType(val firstTypes: List<PyType?>, val paramSpec: PyParamSpecType): PyType {
override fun resolveMember(name: String,
location: PyExpression?,
direction: AccessDirection,
resolveContext: PyResolveContext): List<RatedResolveResult>? {
return null
}
override fun getCompletionVariants(completionPrefix: String?, location: PsiElement?, context: ProcessingContext?): Array<Any> =
emptyArray()
class PyConcatenateType(val firstTypes: List<PyType?>, val paramSpec: PyParamSpecType) : PyCallableParameterVariadicType {
override fun getName(): String = "Concatenate(${firstTypes.joinToString { it?.name ?: "Any" }}, ${paramSpec.name})"
override fun isBuiltin(): Boolean = true
override fun assertValid(message: String?) {
}
}
@@ -1,82 +1,54 @@
package com.jetbrains.python.psi.types;
import com.intellij.openapi.util.text.StringUtil;
import com.intellij.psi.PsiElement;
import com.intellij.util.ArrayUtilRt;
import com.intellij.util.ProcessingContext;
import com.intellij.util.containers.ContainerUtil;
import com.jetbrains.python.PyNames;
import com.jetbrains.python.psi.AccessDirection;
import com.jetbrains.python.psi.PyExpression;
import com.jetbrains.python.psi.PyQualifiedNameOwner;
import com.jetbrains.python.psi.resolve.PyResolveContext;
import com.jetbrains.python.psi.resolve.RatedResolveResult;
import org.jetbrains.annotations.NotNull;
import org.jetbrains.annotations.Nullable;
import java.util.List;
import java.util.Objects;
/**
* Type of typing.ParamSpec using in type checker to unify parameters of generic calls
* Represents a type parameter substituted with a parameter list of a callable as described in
* <a href="https://peps.python.org/pep-0612/">PEP 612 – Parameter Specification Variables</a>.
* <p>
* Declared with either {@code typing.ParamSpec} instantiation or {@code **P} syntax.
* Concrete instantiations of such a parameter are either {@link PyCallableParameterListType} or {@link PyConcatenateType}.
*
* @see PyCallableParameterListType
* @see PyConcatenateType
*/
public final class PyParamSpecType implements PyTypeParameterType {
public final class PyParamSpecType implements PyTypeParameterType, PyCallableParameterVariadicType {
@NotNull private final String myName;
@Nullable private final PyQualifiedNameOwner myDeclarationElement;
@Nullable private final List<PyCallableParameter> myParameters;
@Nullable private final PyType myDefaultType;
@Nullable private final PyCallableParameterVariadicType myDefaultType;
@Nullable private final PyQualifiedNameOwner myScopeOwner;
public PyParamSpecType(@NotNull String name) {
this(name, null, null, null, null);
this(name, null, null, null);
}
private PyParamSpecType(@NotNull String name,
@Nullable PyQualifiedNameOwner declarationElement,
@Nullable List<PyCallableParameter> parameters,
@Nullable PyType defaultType,
@Nullable PyCallableParameterVariadicType defaultType,
@Nullable PyQualifiedNameOwner scopeOwner) {
myName = name;
myDeclarationElement = declarationElement;
myParameters = parameters;
myDefaultType = defaultType;
myScopeOwner = scopeOwner;
}
@NotNull
public PyParamSpecType withParameters(@Nullable List<PyCallableParameter> parameters, @NotNull TypeEvalContext context) {
return new PyParamSpecType(myName, myDeclarationElement, getNonPsiParameters(parameters, context), myDefaultType, myScopeOwner);
}
@NotNull
public PyParamSpecType withDeclarationElement(@Nullable PyQualifiedNameOwner declarationElement) {
return new PyParamSpecType(myName, declarationElement, myParameters, myDefaultType, myScopeOwner);
return new PyParamSpecType(myName, declarationElement, myDefaultType, myScopeOwner);
}
@NotNull
public PyParamSpecType withScopeOwner(@Nullable PyQualifiedNameOwner scopeOwner) {
return new PyParamSpecType(myName, myDeclarationElement, myParameters, myDefaultType, scopeOwner);
return new PyParamSpecType(myName, myDeclarationElement, myDefaultType, scopeOwner);
}
@NotNull
public PyParamSpecType withDefaultType(@Nullable PyType defaultType) {
return new PyParamSpecType(myName, myDeclarationElement, myParameters, defaultType, myScopeOwner);
}
@Nullable
private static List<PyCallableParameter> getNonPsiParameters(@Nullable List<PyCallableParameter> parameters,
@NotNull TypeEvalContext context) {
if (parameters == null) return null;
return ContainerUtil.map(parameters, it -> {
if (it.isPositionalContainer()) return PyCallableParameterImpl.positionalNonPsi(it.getName(), it.getType(context));
if (it.isKeywordContainer()) return PyCallableParameterImpl.keywordNonPsi(it.getName(), it.getType(context));
return PyCallableParameterImpl.nonPsi(it.getName(), it.getType(context));
});
}
@Nullable
public List<PyCallableParameter> getParameters() {
return myParameters;
public PyParamSpecType withDefaultType(@Nullable PyCallableParameterVariadicType defaultType) {
return new PyParamSpecType(myName, myDeclarationElement, defaultType, myScopeOwner);
}
@Override
@@ -84,38 +56,16 @@ public final class PyParamSpecType implements PyTypeParameterType {
return myDeclarationElement;
}
@Nullable
@Override
public List<? extends RatedResolveResult> resolveMember(@NotNull String name,
@Nullable PyExpression location,
@NotNull AccessDirection direction,
@NotNull PyResolveContext resolveContext) {
return null;
}
@Override
public Object[] getCompletionVariants(String completionPrefix, PsiElement location, ProcessingContext context) {
return ArrayUtilRt.EMPTY_OBJECT_ARRAY;
}
@NotNull
@Override
public String getName() {
if (myParameters == null) {
return String.format("ParamSpec(\"%s\")", myName);
}
else {
final TypeEvalContext context = TypeEvalContext.codeInsightFallback(null);
return String.format("[%s]",
StringUtil.join(myParameters, param -> {
if (param != null) {
final PyType type = param.getType(context);
return type != null ? type.getName() : PyNames.UNKNOWN_TYPE;
}
return PyNames.UNKNOWN_TYPE;
},
", "));
}
return "ParamSpec(\"" + myName + "\")";
}
@Override
public String toString() {
String scopeName = myScopeOwner != null ? Objects.requireNonNullElse(myScopeOwner.getQualifiedName(), myScopeOwner.getName()) : null;
return "PyParamSpecType: " + (scopeName != null ? scopeName + ":" : "") + myName;
}
@Override
@@ -125,7 +75,7 @@ public final class PyParamSpecType implements PyTypeParameterType {
@Override
@Nullable
public PyType getDefaultType() {
public PyCallableParameterVariadicType getDefaultType() {
return myDefaultType;
}
@@ -134,15 +84,6 @@ public final class PyParamSpecType implements PyTypeParameterType {
return myName;
}
@Override
public boolean isBuiltin() {
return false;
}
@Override
public void assertValid(String message) {
}
@Override
public boolean equals(Object o) {
if (this == o) {
@@ -152,7 +93,7 @@ public final class PyParamSpecType implements PyTypeParameterType {
return false;
}
final PyParamSpecType type = (PyParamSpecType)o;
return myName.equals(type.myName) && Objects.equals(myScopeOwner, type.myScopeOwner) && Objects.equals(myParameters, type.myParameters);
return myName.equals(type.myName) && Objects.equals(myScopeOwner, type.myScopeOwner);
}
@Override
@@ -363,10 +363,8 @@ public final class PyTypeChecker {
private static boolean match(@NotNull PyParamSpecType expected, @Nullable PyType actual, @NotNull MatchContext context) {
if (actual == null) return true;
if (!(actual instanceof PyParamSpecType callableActual)) return false;
final var parameters = callableActual.getParameters();
if (parameters == null) return false;
context.mySubstitutions.paramSpecs.put(expected, expected.withParameters(parameters, context.context));
if (!(actual instanceof PyCallableParameterListType actualParameters)) return false;
context.mySubstitutions.paramSpecs.put(expected, actualParameters);
return true;
}
@@ -612,7 +610,7 @@ public final class PyTypeChecker {
final var firstExpectedParam = expectedParameters.get(0);
final var expectedParamType = firstExpectedParam.getType(context);
if (expectedParamType instanceof final PyParamSpecType expectedParamSpecType) {
matchContext.mySubstitutions.paramSpecs.put(expectedParamSpecType, expectedParamSpecType.withParameters(actualParameters, context));
matchContext.mySubstitutions.paramSpecs.put(expectedParamSpecType, new PyCallableParameterListTypeImpl(actualParameters));
return true;
}
else if (expectedParamType instanceof final PyConcatenateType expectedConcatenateType) {
@@ -640,8 +638,7 @@ public final class PyTypeChecker {
if (actualParamRightBound < actualParameters.size()) {
final var expectedParamSpecType = expectedConcatenateType.getParamSpec();
final var restActualParameters = actualParameters.subList(actualParamRightBound, actualParameters.size());
final var parametersSubst = expectedParamSpecType.withParameters(restActualParameters, context);
matchContext.mySubstitutions.paramSpecs.put(expectedParamSpecType, parametersSubst);
matchContext.mySubstitutions.paramSpecs.put(expectedParamSpecType, new PyCallableParameterListTypeImpl(restActualParameters));
return true;
}
}
@@ -794,8 +791,8 @@ public final class PyTypeChecker {
result.typeVars.put(typeVarType, entry.getValue());
}
else if (entry.getKey() instanceof PyTypeVarTupleType typeVarTuple) {
assert entry.getValue() instanceof PyVariadicType;
result.typeVarTuples.put(typeVarTuple, (PyVariadicType)entry.getValue());
assert entry.getValue() instanceof PyPositionalVariadicType;
result.typeVarTuples.put(typeVarTuple, (PyPositionalVariadicType)entry.getValue());
}
// TODO Handle ParamSpecs here
}
@@ -953,23 +950,18 @@ public final class PyTypeChecker {
}
}
for (PyParamSpecType paramSpecType : typeParamsFromReturnType.paramSpecs) {
// ParamSpecs already bound to parameter lists
if (paramSpecType.getParameters() != null) {
continue;
}
boolean canGetBoundFromArguments = typeParamsFromParameterTypes.paramSpecs.contains(paramSpecType);
boolean isAlreadyBound = existingSubstitutions.paramSpecs.containsKey(paramSpecType);
if (canGetBoundFromArguments && !isAlreadyBound) {
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));
PyCallableParameterVariadicType defaultType = paramSpecType.getDefaultType();
existingSubstitutions.paramSpecs.put(paramSpecType, defaultType);
}
else {
existingSubstitutions.paramSpecs.put(paramSpecType, new PyCallableParameterListTypeImpl(
List.of(PyCallableParameterImpl.positionalNonPsi("args", null),
PyCallableParameterImpl.keywordNonPsi("kwargs", null)))
);
}
}
}
@@ -1010,7 +1002,6 @@ public final class PyTypeChecker {
if (type instanceof PyTypeVarTupleType typeVarTupleType) {
generics.typeVarTuples.add(typeVarTupleType);
}
// TODO Filter out PyParamSpecTypes representing actual lists of parameters, not type parameters declared via ParamSpec
if (type instanceof PyParamSpecType) {
generics.paramSpecs.add((PyParamSpecType)type);
}
@@ -1054,6 +1045,11 @@ public final class PyTypeChecker {
}
collectGenerics(concatenateType.getParamSpec(), context, generics, visited);
}
else if (type instanceof PyCallableParameterListType callableParameterList) {
for (PyCallableParameter parameter : callableParameterList.getParameters()) {
collectGenerics(parameter.getType(context), context, generics, visited);
}
}
else if (type instanceof PyUnpackedTupleType unpackedTupleType) {
for (PyType elementType : unpackedTupleType.getElementTypes()) {
collectGenerics(elementType, context, generics, visited);
@@ -1162,7 +1158,7 @@ public final class PyTypeChecker {
}
return substitution;
}
else if (type instanceof PyParamSpecType paramSpecType && paramSpecType.getParameters() == null) {
else if (type instanceof PyParamSpecType paramSpecType) {
if (!substitutions.paramSpecs.containsKey(paramSpecType)) {
PyParamSpecType sameScopeSubstitution = StreamEx.of(substitutions.paramSpecs.keySet())
.findFirst(typeVarType -> {
@@ -1175,7 +1171,7 @@ public final class PyTypeChecker {
}
return paramSpecType;
}
PyParamSpecType substitution = substitutions.paramSpecs.get(paramSpecType);
PyCallableParameterVariadicType substitution = substitutions.paramSpecs.get(paramSpecType);
if (substitution != null && !substitution.equals(paramSpecType) && hasGenerics(substitution, context)) {
return substitute(substitution, substitutions, context, substituting);
}
@@ -1235,15 +1231,16 @@ public final class PyTypeChecker {
PyCallableParameter onlyParam = ContainerUtil.getOnlyItem(parameters);
if (onlyParam != null && onlyParam.getType(context) instanceof PyParamSpecType paramSpecType) {
final var substitution = substitute(paramSpecType, substitutions, context);
if (substitution instanceof PyParamSpecType paramSpecSubst && paramSpecSubst.getParameters() != null) {
substParams = paramSpecSubst.getParameters();
if (substitution instanceof PyCallableParameterListType callableParams) {
substParams = callableParams.getParameters();
}
else {
substParams = List.of(PyCallableParameterImpl.nonPsi(substitution));
}
}
else if (onlyParam != null && onlyParam.getType(context) instanceof PyConcatenateType concatenateType) {
final var paramSpecSubst = as(substitute(concatenateType.getParamSpec(), substitutions, context), PyParamSpecType.class);
PyCallableParameterListType paramSpecSubst =
as(substitute(concatenateType.getParamSpec(), substitutions, context), PyCallableParameterListType.class);
List<PyCallableParameter> paramSpecParams = paramSpecSubst != null ? paramSpecSubst.getParameters() : null;
substParams = StreamEx.of(concatenateType.getFirstTypes())
.flatCollection(paramType -> substituteExpand(paramType, substitutions, context, substituting))
@@ -1447,7 +1444,7 @@ public final class PyTypeChecker {
for (Map.Entry<PyTypeVarTupleType, PyPositionalVariadicType> typeVarMapping : newSubstitutions.typeVarTuples.entrySet()) {
substitutions.typeVarTuples.putIfAbsent(typeVarMapping.getKey(), typeVarMapping.getValue());
}
for (Map.Entry<PyParamSpecType, PyParamSpecType> paramSpecMapping : newSubstitutions.paramSpecs.entrySet()) {
for (Map.Entry<PyParamSpecType, PyCallableParameterVariadicType> paramSpecMapping : newSubstitutions.paramSpecs.entrySet()) {
substitutions.paramSpecs.putIfAbsent(paramSpecMapping.getKey(), paramSpecMapping.getValue());
}
});
@@ -1641,10 +1638,10 @@ public final class PyTypeChecker {
substitutions.typeVars.put(typeVar, pair.getSecond());
}
else if (pair.getFirst() instanceof PyTypeVarTupleType typeVarTuple) {
substitutions.typeVarTuples.put(typeVarTuple, as(pair.getSecond(), PyVariadicType.class));
substitutions.typeVarTuples.put(typeVarTuple, as(pair.getSecond(), PyPositionalVariadicType.class));
}
else if (pair.getFirst() instanceof PyParamSpecType paramSpec) {
substitutions.paramSpecs.put(paramSpec, as(pair.getSecond(), PyParamSpecType.class));
substitutions.paramSpecs.put(paramSpec, as(pair.getSecond(), PyCallableParameterVariadicType.class));
}
}
return substitutions;
@@ -1711,7 +1708,7 @@ public final class PyTypeChecker {
private final Map<PyTypeVarTupleType, PyPositionalVariadicType> typeVarTuples;
@NotNull
private final Map<PyParamSpecType, PyParamSpecType> paramSpecs;
private final Map<PyParamSpecType, PyCallableParameterVariadicType> paramSpecs;
@Nullable
private PyType qualifierType;
@@ -1727,7 +1724,7 @@ public final class PyTypeChecker {
.toCustomMap(LinkedHashMap::new),
EntryStream.of(typeParameters)
.selectKeys(PyParamSpecType.class)
.selectValues(PyParamSpecType.class)
.selectValues(PyCallableParameterVariadicType.class)
.toCustomMap(LinkedHashMap::new),
null
);
@@ -1739,7 +1736,7 @@ public final class PyTypeChecker {
private GenericSubstitutions(@NotNull Map<PyTypeVarType, PyType> typeVars,
@NotNull Map<PyTypeVarTupleType, PyPositionalVariadicType> typeVarTuples,
@NotNull Map<PyParamSpecType, PyParamSpecType> paramSpecs,
@NotNull Map<PyParamSpecType, PyCallableParameterVariadicType> paramSpecs,
@Nullable PyType qualifierType) {
this.typeVars = typeVars;
this.typeVarTuples = typeVarTuples;
@@ -1747,7 +1744,7 @@ public final class PyTypeChecker {
this.qualifierType = qualifierType;
}
public @NotNull Map<PyParamSpecType, PyParamSpecType> getParamSpecs() {
public @NotNull Map<PyParamSpecType, PyCallableParameterVariadicType> getParamSpecs() {
return Collections.unmodifiableMap(paramSpecs);
}
@@ -20,10 +20,19 @@ public final class PyTypeParameterMapping {
for (Couple<PyType> couple : mapping) {
PyType expectedType = couple.getFirst();
PyType actualType = couple.getSecond();
if (expectedType instanceof PyPositionalVariadicType && !(actualType instanceof PyPositionalVariadicType || actualType == null)) {
throw new IllegalArgumentException("Variadic type " + expectedType + " cannot be mapped to a non-variadic type " + actualType);
if (expectedType instanceof PyPositionalVariadicType &&
!(actualType instanceof PyPositionalVariadicType || actualType == null)) {
throw new IllegalArgumentException(
"Positional variadic type " + expectedType + " cannot be mapped to a non-variadic type " + actualType
);
}
if (!(expectedType instanceof PyPositionalVariadicType) && actualType instanceof PyPositionalVariadicType) {
if (expectedType instanceof PyCallableParameterVariadicType &&
!(actualType instanceof PyCallableParameterVariadicType || actualType == null)) {
throw new IllegalArgumentException(
"Callable parameter variadic type " + expectedType + " cannot be mapped to a non-variadic type " + actualType
);
}
if (!(expectedType instanceof PyVariadicType) && actualType instanceof PyVariadicType) {
throw new IllegalArgumentException("Non-variadic type " + expectedType + " cannot be mapped to a variadic type " + actualType);
}
}
@@ -188,10 +197,8 @@ public final class PyTypeParameterMapping {
else if (expectedTypesDeque.size() == 1) {
PyType onlyLeftExpectedType = expectedTypesDeque.peekFirst();
if (onlyLeftExpectedType instanceof PyPositionalVariadicType) {
if (actualTypesDeque.size() == 1 && actualTypesDeque.peekFirst() instanceof PyPositionalVariadicType variadicType) {
if (onlyLeftExpectedType instanceof PyVariadicType) {
// [*Ts] <- [*Ts] or [*Ts] <- [*tuple[T1, ...]]
if (actualTypesDeque.size() == 1 && actualTypesDeque.peekFirst() instanceof PyVariadicType variadicType) {
if (actualTypesDeque.size() == 1 && actualTypesDeque.peekFirst() instanceof PyPositionalVariadicType variadicType) {
centerMappedTypes.add(Couple.of(onlyLeftExpectedType, variadicType));
}
else {