mirror of
https://gitflic.ru/project/openide/openide.git
synced 2026-09-27 10:03:11 +07:00
PY-76883 Implement proper matching for callable signatures
Introduce `PyCallableParameterMapping`, a new utility responsible for matching Python callable signatures. Spec: https://typing.python.org/en/latest/spec/callables.html#assignability-rules-for-callables Key features: - Support for complex parameter kinds: positional-only, keyword-only, variadic (*args, **kwargs), and TypeVarTuples. - Implementation of assignability rules for callables, ensuring that actual parameters correctly satisfy expected signatures. - Handling of parameter defaults. - Comprehensive categorization of `PyCallableParameter` lists based on Python's parameter syntax separators (/ and *). This utility facilitates more accurate type checking for higher-order functions and Protocol implementations involving complex signatures. The old `PyTypeParameterMapping#mapByShape` is sunsetted in favor of the new logic mentioned above. Co-authored-by: Mikhail Golubev <mikhail.golubev@jetbrains.com> Co-authored-by: Petr Golubev <petr.golubev@jetbrains.com> GitOrigin-RevId: 9b4ed6445bf7ad6af1409e64c31750210be6b142
This commit is contained in:
committed by
intellij-monorepo-bot
co-authored by
Mikhail Golubev
Petr Golubev
parent
6c42ce2770
commit
e2600feb97
+349
@@ -0,0 +1,349 @@
|
||||
package com.jetbrains.python.psi.types
|
||||
|
||||
import com.jetbrains.python.isPrivate
|
||||
import org.jetbrains.annotations.ApiStatus
|
||||
import java.util.*
|
||||
|
||||
/**
|
||||
* Matches signatures of callables according to the Callable compatibility rules
|
||||
* @see <a href="https://typing.python.org/en/latest/spec/callables.html#assignability-rules-for-callables">Specification</a>
|
||||
*/
|
||||
@ApiStatus.Internal
|
||||
object PyCallableParameterMapping {
|
||||
|
||||
private enum class ParameterKind {
|
||||
POSITIONAL_ONLY,
|
||||
POSITIONAL_OR_KEYWORD,
|
||||
KEYWORD_ONLY,
|
||||
POSITIONAL_CONTAINER,
|
||||
KEYWORD_CONTAINER,
|
||||
TYPE_VAR_TUPLE;
|
||||
}
|
||||
|
||||
private data class Parameter(
|
||||
val parameter: PyCallableParameter,
|
||||
val kind: ParameterKind,
|
||||
) {
|
||||
val name: String? get() = parameter.name
|
||||
val hasDefault: Boolean get() = parameter.hasDefaultValue()
|
||||
|
||||
fun getArgumentType(context: TypeEvalContext): PyType? = parameter.getArgumentType(context)
|
||||
|
||||
val acceptsPositionalArgument =
|
||||
kind == ParameterKind.POSITIONAL_ONLY || kind == ParameterKind.POSITIONAL_OR_KEYWORD || kind == ParameterKind.POSITIONAL_CONTAINER
|
||||
|
||||
val acceptsKeywordArgument =
|
||||
kind == ParameterKind.KEYWORD_ONLY || kind == ParameterKind.POSITIONAL_OR_KEYWORD || kind == ParameterKind.KEYWORD_CONTAINER
|
||||
}
|
||||
|
||||
/**
|
||||
* Maps a list of expected parameters to a list of actual parameters for a Callable
|
||||
* by analyzing their signatures and attempting to align them.
|
||||
*
|
||||
* @param expectedCallableParameters a list of callable parameters from the expected Callable.
|
||||
* @param actualCallableParameters a list of callable parameters from the actual Callable.
|
||||
* @param context TypeEvalContext.
|
||||
* @return a `PyCallableParameterMapping` object representing the result of the parameter mapping,
|
||||
* or `null` if one of the signatures is specified incorrectly, or they do not match by structure
|
||||
*/
|
||||
@JvmStatic
|
||||
fun mapCallableParameters(
|
||||
expectedCallableParameters: List<PyCallableParameter>,
|
||||
actualCallableParameters: List<PyCallableParameter>,
|
||||
context: TypeEvalContext,
|
||||
): PyTypeParameterMapping? {
|
||||
val expectedCategorizedParameters = categorizeParameters(expectedCallableParameters, context) ?: return null
|
||||
val actualCategorizedParameters = categorizeParameters(actualCallableParameters, context) ?: return null
|
||||
|
||||
val expectedParameters = ArrayDeque(expectedCategorizedParameters)
|
||||
val actualParameters = ArrayDeque(actualCategorizedParameters)
|
||||
|
||||
val expectedTypes = mutableListOf<PyType?>()
|
||||
val actualTypes = mutableListOf<PyType?>()
|
||||
|
||||
val actualKeywordContainer = actualParameters.firstOrNull { it.kind == ParameterKind.KEYWORD_CONTAINER }
|
||||
|
||||
val actualHasTypeVarTuple = actualParameters.any { it.kind == ParameterKind.TYPE_VAR_TUPLE }
|
||||
|
||||
val actualKeywordOrPositional = actualParameters
|
||||
.filter {
|
||||
(it.kind == ParameterKind.POSITIONAL_OR_KEYWORD ||
|
||||
it.kind == ParameterKind.KEYWORD_ONLY) && it.name != null
|
||||
}
|
||||
.associateBy { it.name!! }
|
||||
.toMutableMap()
|
||||
|
||||
/**
|
||||
* if Callable contains TypeVarTuple, say, Callable[[int, str, *Ts, bool, int], None]
|
||||
* we need to calculate how many *positional-only* parameters
|
||||
* (positional including positional container in case of real signature)
|
||||
* are expected from the right side of the TypeVarTuple.
|
||||
*/
|
||||
val actualPositionalOrContainerCount = actualParameters.count { it.acceptsPositionalArgument }
|
||||
val expectedPositionalOrContainerCount = expectedParameters.count { it.acceptsPositionalArgument }
|
||||
|
||||
val paramsTypeVarTupleShouldAccept =
|
||||
actualPositionalOrContainerCount - expectedPositionalOrContainerCount + (if (actualHasTypeVarTuple) 1 else 0)
|
||||
|
||||
while (expectedParameters.isNotEmpty() && actualParameters.isNotEmpty()) {
|
||||
val expectedParameter = expectedParameters.peek()
|
||||
val actualParameter = actualParameters.peek()
|
||||
|
||||
when (expectedParameter.kind) {
|
||||
// Positional-only can match with positional or positional-or-keyword parameters
|
||||
ParameterKind.POSITIONAL_ONLY -> {
|
||||
if (actualParameter.parameter.isPositionalContainer) {
|
||||
expectedTypes.add(expectedParameter.getArgumentType(context))
|
||||
actualTypes.add(actualParameter.getArgumentType(context))
|
||||
expectedParameters.pop()
|
||||
}
|
||||
else if (actualParameter.acceptsPositionalArgument) {
|
||||
if (expectedParameter.hasDefault && !actualParameter.hasDefault) {
|
||||
return null
|
||||
}
|
||||
expectedTypes.add(expectedParameter.getArgumentType(context))
|
||||
expectedParameters.pop()
|
||||
actualTypes.add(actualParameter.getArgumentType(context))
|
||||
actualParameters.pop()
|
||||
}
|
||||
else {
|
||||
return null
|
||||
}
|
||||
}
|
||||
// Must match by name
|
||||
ParameterKind.POSITIONAL_OR_KEYWORD -> {
|
||||
if (actualParameter.kind == ParameterKind.POSITIONAL_OR_KEYWORD) {
|
||||
when {
|
||||
expectedParameter.parameter.isSelf && actualParameter.parameter.isSelf -> continue
|
||||
expectedParameter.name != actualParameter.name -> return null
|
||||
expectedParameter.hasDefault && !actualParameter.hasDefault -> return null
|
||||
}
|
||||
expectedTypes.add(expectedParameter.getArgumentType(context))
|
||||
expectedParameters.pop()
|
||||
actualTypes.add(actualParameter.getArgumentType(context))
|
||||
actualParameters.pop()
|
||||
}
|
||||
else if (actualParameter.kind == ParameterKind.POSITIONAL_CONTAINER && actualKeywordContainer != null) {
|
||||
expectedTypes.add(expectedParameter.getArgumentType(context))
|
||||
actualTypes.add(actualParameter.getArgumentType(context))
|
||||
expectedParameters.pop()
|
||||
}
|
||||
else {
|
||||
return null
|
||||
}
|
||||
}
|
||||
// Keyword-only parameter must match by name from the set of keyword-only or
|
||||
// positional-or-keyword parameters from the actual signature
|
||||
ParameterKind.KEYWORD_ONLY -> {
|
||||
val actualKwOnlyOrPositionalParam = actualKeywordOrPositional.remove(expectedParameter.name)
|
||||
if (actualKwOnlyOrPositionalParam != null) {
|
||||
if (expectedParameter.hasDefault && !actualKwOnlyOrPositionalParam.hasDefault) {
|
||||
return null
|
||||
}
|
||||
expectedTypes.add(expectedParameter.getArgumentType(context))
|
||||
expectedParameters.pop()
|
||||
actualTypes.add(actualKwOnlyOrPositionalParam.getArgumentType(context))
|
||||
require(actualParameters.remove(actualKwOnlyOrPositionalParam))
|
||||
}
|
||||
else if (actualKeywordContainer != null) {
|
||||
expectedTypes.add(expectedParameter.getArgumentType(context))
|
||||
actualTypes.add(actualKeywordContainer.getArgumentType(context))
|
||||
// All keyword parameters and the corresponding container are consumed, so we can pop the container
|
||||
expectedParameters.pop()
|
||||
}
|
||||
else {
|
||||
return null
|
||||
}
|
||||
}
|
||||
// *args can consume multiple positional parameters
|
||||
ParameterKind.POSITIONAL_CONTAINER -> {
|
||||
val argsType = expectedParameter.getArgumentType(context)
|
||||
// Consume all remaining positional parameters
|
||||
var actualPositional = actualParameters.pop()
|
||||
while (actualParameters.isNotEmpty() && !actualPositional.parameter.isPositionalContainer) {
|
||||
if (!actualPositional.acceptsPositionalArgument || !actualPositional.hasDefault) {
|
||||
return null
|
||||
}
|
||||
expectedTypes.add(argsType)
|
||||
actualTypes.add(actualPositional.getArgumentType(context))
|
||||
actualPositional = actualParameters.pop()
|
||||
}
|
||||
// we need to match containers themselves as well (e.g. *args: T <- *args: T1)
|
||||
if (actualPositional.parameter.isPositionalContainer) {
|
||||
// *args: T is not propagated to *tuple[T, ...] here to avoid conflicts with mapping of subsequent types
|
||||
expectedTypes.add(argsType)
|
||||
actualTypes.add(actualPositional.getArgumentType(context))
|
||||
// All positional parameters and the corresponding container are consumed, so we can pop the container
|
||||
expectedParameters.pop()
|
||||
}
|
||||
}
|
||||
// **kwargs can consume multiple keyword parameters
|
||||
ParameterKind.KEYWORD_CONTAINER -> {
|
||||
val kwargsType = expectedParameter.getArgumentType(context)
|
||||
// Consume all remaining keyword parameters
|
||||
var actualKeywordParam = actualParameters.pop()
|
||||
while (actualParameters.isNotEmpty() && !actualKeywordParam.parameter.isKeywordContainer) {
|
||||
if (!(actualKeywordParam.acceptsKeywordArgument) || !actualKeywordParam.hasDefault) {
|
||||
return null
|
||||
}
|
||||
expectedTypes.add(kwargsType)
|
||||
actualTypes.add(actualKeywordParam.getArgumentType(context))
|
||||
|
||||
actualKeywordParam = actualParameters.pop()
|
||||
}
|
||||
// match keyword containers themselves
|
||||
if (actualKeywordParam.parameter.isKeywordContainer) {
|
||||
expectedTypes.add(kwargsType)
|
||||
actualTypes.add(actualKeywordParam.getArgumentType(context))
|
||||
expectedParameters.pop()
|
||||
}
|
||||
}
|
||||
// TypeVarTuple (*Ts) consumes a calculated number of positional parameters
|
||||
ParameterKind.TYPE_VAR_TUPLE -> {
|
||||
expectedTypes.add(expectedParameter.getArgumentType(context))
|
||||
|
||||
repeat(paramsTypeVarTupleShouldAccept) {
|
||||
if (actualParameters.isEmpty()) return null
|
||||
val actualParameter = actualParameters.pop()
|
||||
if (!actualParameter.acceptsPositionalArgument && actualParameter.kind != ParameterKind.TYPE_VAR_TUPLE) {
|
||||
return null
|
||||
}
|
||||
var actualType = actualParameter.getArgumentType(context)
|
||||
// *args: T mapped to TypeVarTuple should be represented as *tuple[T, ...]
|
||||
if (actualParameter.parameter.isPositionalContainer && actualType !is PyPositionalVariadicType) {
|
||||
actualType = PyUnpackedTupleTypeImpl.createUnbound(actualType)
|
||||
}
|
||||
actualTypes.add(actualType)
|
||||
}
|
||||
// TypeVarTuple consumed all the required parameters.
|
||||
expectedParameters.pop()
|
||||
}
|
||||
}
|
||||
}
|
||||
// Post-process remaining params
|
||||
while (expectedParameters.isNotEmpty()) {
|
||||
val parameter = expectedParameters.pop()
|
||||
val type = parameter.getArgumentType(context)
|
||||
when {
|
||||
// TODO remove these container checks when type parameters in protocols are properly substituted
|
||||
parameter.parameter.isPositionalContainer && (type is PyPositionalVariadicType || type is PyParamSpecType) -> continue
|
||||
parameter.parameter.isKeywordContainer && type is PyParamSpecType -> continue
|
||||
parameter.kind == ParameterKind.TYPE_VAR_TUPLE -> { // can be empty
|
||||
if (actualTypes.isEmpty()) {
|
||||
expectedTypes.add(type)
|
||||
}
|
||||
}
|
||||
else -> return null
|
||||
}
|
||||
}
|
||||
// Handle unmatched actual parameters (extra in actual)
|
||||
while (actualParameters.isNotEmpty()) {
|
||||
val parameter = actualParameters.pop()
|
||||
when {
|
||||
parameter.parameter.isPositionalContainer || parameter.parameter.isKeywordContainer -> continue
|
||||
parameter.hasDefault -> continue
|
||||
else -> return null
|
||||
}
|
||||
}
|
||||
|
||||
return PyTypeParameterMapping.mapByShape(expectedTypes, actualTypes)
|
||||
}
|
||||
|
||||
// Unwraps *args: *tuple[T, T1] to a sequence of positional parameters: T, T1 for matching purposes
|
||||
private fun flattenPositionalContainer(parameters: List<Parameter>, context: TypeEvalContext): List<Parameter> {
|
||||
val flattenedParameters = mutableListOf<Parameter>()
|
||||
for (parameter in parameters) {
|
||||
val parameterType = parameter.getArgumentType(context)
|
||||
if (parameter.parameter.isPositionalContainer && parameterType is PyUnpackedTupleType && !parameterType.isUnbound()) {
|
||||
val flattenedTypes = parameterType.getElementTypes()
|
||||
val unwrappedArgParameters = flattenedTypes.map { type ->
|
||||
Parameter(PyCallableParameterImpl.nonPsi(type), ParameterKind.POSITIONAL_ONLY)
|
||||
}
|
||||
flattenedParameters.addAll(unwrappedArgParameters)
|
||||
}
|
||||
else {
|
||||
flattenedParameters.add(parameter)
|
||||
}
|
||||
}
|
||||
return flattenedParameters
|
||||
}
|
||||
|
||||
private enum class ParameterState {
|
||||
POSITIONAL_OR_KEYWORD,
|
||||
KEYWORD_ONLY,
|
||||
POSITIONAL_CONTAINER,
|
||||
}
|
||||
|
||||
/**
|
||||
* Categorizes parameters into based on Python's parameter syntax.
|
||||
* Returns null if the parameter list is invalid (e.g., duplicate *args or **kwargs).
|
||||
*/
|
||||
private fun categorizeParameters(callableParameters: List<PyCallableParameter>, context: TypeEvalContext): List<Parameter>? {
|
||||
val parameters = mutableListOf<Parameter>()
|
||||
|
||||
var positionalContainer: Parameter? = null
|
||||
var keywordContainer: Parameter? = null
|
||||
var typeVarTupleType: PyTypeVarTupleType? = null
|
||||
|
||||
var state = ParameterState.POSITIONAL_OR_KEYWORD
|
||||
for (param in callableParameters) {
|
||||
val argumentType = param.getArgumentType(context)
|
||||
if (param.isPositionOnlySeparator) {
|
||||
parameters.replaceAll { p -> p.copy(kind = ParameterKind.POSITIONAL_ONLY) }
|
||||
continue
|
||||
}
|
||||
else if (param.isKeywordContainer) {
|
||||
if (keywordContainer != null) return null
|
||||
|
||||
keywordContainer = Parameter(param, ParameterKind.KEYWORD_CONTAINER)
|
||||
parameters.add(keywordContainer)
|
||||
continue
|
||||
}
|
||||
else if (param.isPositionalContainer) {
|
||||
if (positionalContainer != null) return null
|
||||
|
||||
positionalContainer = Parameter(param, ParameterKind.POSITIONAL_CONTAINER)
|
||||
parameters.add(positionalContainer)
|
||||
state = ParameterState.POSITIONAL_CONTAINER
|
||||
continue
|
||||
}
|
||||
if (argumentType is PyTypeVarTupleType) {
|
||||
if (typeVarTupleType != null) return null // Only one TypeVarTuple is allowed
|
||||
|
||||
typeVarTupleType = argumentType
|
||||
parameters.add(Parameter(param, ParameterKind.TYPE_VAR_TUPLE))
|
||||
continue
|
||||
}
|
||||
else if (param.isKeywordOnlySeparator) {
|
||||
state = ParameterState.KEYWORD_ONLY
|
||||
continue
|
||||
}
|
||||
else {
|
||||
if (state == ParameterState.POSITIONAL_OR_KEYWORD) {
|
||||
val paramName = param.name
|
||||
val isPositionalOnly = paramName == null || isPrivate(paramName)
|
||||
if (isPositionalOnly) {
|
||||
parameters.lastOrNull()?.let {
|
||||
if (it.kind != ParameterKind.POSITIONAL_ONLY && it.kind != ParameterKind.TYPE_VAR_TUPLE) {
|
||||
// Previous parameter should also be positional-only in this case
|
||||
return null
|
||||
}
|
||||
}
|
||||
parameters.add(Parameter(param, ParameterKind.POSITIONAL_ONLY))
|
||||
}
|
||||
else {
|
||||
parameters.add(Parameter(param, ParameterKind.POSITIONAL_OR_KEYWORD))
|
||||
}
|
||||
}
|
||||
else {
|
||||
val name = param.name
|
||||
if (name == null) {
|
||||
return null
|
||||
}
|
||||
parameters.add(Parameter(param, ParameterKind.KEYWORD_ONLY))
|
||||
}
|
||||
}
|
||||
}
|
||||
return flattenPositionalContainer(parameters, context)
|
||||
}
|
||||
}
|
||||
@@ -861,36 +861,9 @@ public final class PyTypeChecker {
|
||||
}
|
||||
}
|
||||
|
||||
final var hasParamSpec = ContainerUtil.exists(expectedParameters, it -> (it.getType(context) instanceof PyParamSpecType));
|
||||
|
||||
// TODO Implement proper compatibility check for callable signatures, including positional- and keyword-only arguments, defaults, etc.
|
||||
if (!hasParamSpec) {
|
||||
boolean shouldAcceptUnlimitedPositionalArgs = ContainerUtil.exists(expectedParameters, PyCallableParameter::isPositionalContainer);
|
||||
boolean canAcceptUnlimitedPositionalArgs = ContainerUtil.exists(actualParameters, PyCallableParameter::isPositionalContainer);
|
||||
if (shouldAcceptUnlimitedPositionalArgs && !canAcceptUnlimitedPositionalArgs) return false;
|
||||
|
||||
boolean shouldAcceptArbitraryKeywordArgs = ContainerUtil.exists(expectedParameters, PyCallableParameter::isKeywordContainer);
|
||||
boolean canAcceptArbitraryKeywordArgs = ContainerUtil.exists(actualParameters, PyCallableParameter::isKeywordContainer);
|
||||
if (shouldAcceptArbitraryKeywordArgs && !canAcceptArbitraryKeywordArgs) return false;
|
||||
}
|
||||
|
||||
List<PyType> expectedElementTypes = StreamEx.of(expectedParameters)
|
||||
.filter(cp -> !(cp.getParameter() instanceof PySlashParameter || cp.getParameter() instanceof PySingleStarParameter))
|
||||
.map(cp -> {
|
||||
PyType argType = cp.getArgumentType(context);
|
||||
if (cp.isPositionalContainer() && !(argType instanceof PyPositionalVariadicType)) {
|
||||
return PyUnpackedTupleTypeImpl.createUnbound(argType);
|
||||
}
|
||||
return argType;
|
||||
})
|
||||
.toList();
|
||||
|
||||
final var expectedElementTypes2 =
|
||||
ContainerUtil.filter(expectedElementTypes, type -> !(type instanceof PyParamSpecType));
|
||||
|
||||
PyTypeParameterMapping mapping = PyTypeParameterMapping.mapWithParameterList(ContainerUtil.subList(expectedElementTypes2, startIndex),
|
||||
ContainerUtil.subList(actualParameters, startIndex),
|
||||
context);
|
||||
var mapping = PyCallableParameterMapping.mapCallableParameters(ContainerUtil.subList(expectedParameters, startIndex),
|
||||
ContainerUtil.subList(actualParameters, startIndex),
|
||||
context);
|
||||
if (mapping == null) {
|
||||
return false;
|
||||
}
|
||||
|
||||
@@ -31,82 +31,6 @@ public final class PyTypeParameterMapping {
|
||||
myMappedTypes = mapping;
|
||||
}
|
||||
|
||||
public static @Nullable PyTypeParameterMapping mapWithParameterList(@NotNull List<? extends PyType> expectedParameterTypes,
|
||||
@NotNull List<PyCallableParameter> actualParameters,
|
||||
@NotNull TypeEvalContext context) {
|
||||
List<PyType> flattenedExpectedParameterTypes = flattenUnpackedTupleTypes(expectedParameterTypes);
|
||||
int expectedArity = ContainerUtil.exists(flattenedExpectedParameterTypes, Conditions.instanceOf(PyPositionalVariadicType.class))
|
||||
? -1
|
||||
: flattenedExpectedParameterTypes.size();
|
||||
|
||||
List<PyType> requiredPositionalArgumentTypes = new ArrayList<>();
|
||||
List<PyType> optionalPositionalArgumentTypes = new ArrayList<>();
|
||||
List<PyType> positionalVarargArgumentTypes = new SmartList<>();
|
||||
|
||||
for (PyCallableParameter parameter : actualParameters) {
|
||||
if (parameter.isSelf()
|
||||
|| parameter.getParameter() instanceof PySlashParameter
|
||||
|| parameter.getParameter() instanceof PySingleStarParameter
|
||||
|| parameter.isKeywordContainer()) {
|
||||
continue;
|
||||
}
|
||||
if (parameter.getParameter() instanceof PyNamedParameter namedParameter && namedParameter.isKeywordOnly()) {
|
||||
if (!namedParameter.hasDefaultValue()) {
|
||||
return null;
|
||||
}
|
||||
continue;
|
||||
}
|
||||
PyType actualParameterType = parameter.getType(context);
|
||||
if (parameter.isPositionalContainer()) {
|
||||
// Convert an automatic tuple[MyType, ...] to *tuple[MyType, ...] for matching purposes
|
||||
if (actualParameterType instanceof PyTupleType argsTupleType) {
|
||||
positionalVarargArgumentTypes.addAll(flattenUnpackedTupleTypes(Collections.singletonList(argsTupleType.asUnpackedTupleType())));
|
||||
}
|
||||
else {
|
||||
positionalVarargArgumentTypes.addAll(flattenUnpackedTupleTypes(Collections.singletonList(actualParameterType)));
|
||||
}
|
||||
}
|
||||
else if (parameter.hasDefaultValue()) {
|
||||
optionalPositionalArgumentTypes.add(actualParameterType);
|
||||
}
|
||||
else {
|
||||
requiredPositionalArgumentTypes.addAll(flattenUnpackedTupleTypes(Collections.singletonList(actualParameterType)));
|
||||
}
|
||||
}
|
||||
|
||||
if (positionalVarargArgumentTypes.size() > 1 ||
|
||||
positionalVarargArgumentTypes.size() == 1 && !(positionalVarargArgumentTypes.get(0) instanceof PyPositionalVariadicType)) {
|
||||
requiredPositionalArgumentTypes.addAll(optionalPositionalArgumentTypes);
|
||||
optionalPositionalArgumentTypes.clear();
|
||||
requiredPositionalArgumentTypes.addAll(positionalVarargArgumentTypes);
|
||||
positionalVarargArgumentTypes.clear();
|
||||
}
|
||||
|
||||
int actualArity = ContainerUtil.exists(requiredPositionalArgumentTypes, Conditions.instanceOf(PyPositionalVariadicType.class)) ?
|
||||
-1 :
|
||||
requiredPositionalArgumentTypes.size();
|
||||
|
||||
|
||||
if (expectedArity != -1 && actualArity != -1) {
|
||||
if (actualArity > expectedArity) {
|
||||
return null;
|
||||
}
|
||||
List<PyType> arityAdjustedActualParameterTypes = new ArrayList<>(requiredPositionalArgumentTypes);
|
||||
arityAdjustedActualParameterTypes.addAll(optionalPositionalArgumentTypes.subList(
|
||||
0, Math.min(optionalPositionalArgumentTypes.size(), expectedArity - arityAdjustedActualParameterTypes.size())
|
||||
));
|
||||
if (!positionalVarargArgumentTypes.isEmpty() && expectedArity - arityAdjustedActualParameterTypes.size() > 0) {
|
||||
assert positionalVarargArgumentTypes.size() == 1 && positionalVarargArgumentTypes.get(0) instanceof PyPositionalVariadicType;
|
||||
arityAdjustedActualParameterTypes.add(positionalVarargArgumentTypes.get(0));
|
||||
}
|
||||
return mapByShape(flattenedExpectedParameterTypes, arityAdjustedActualParameterTypes);
|
||||
}
|
||||
return mapByShape(flattenedExpectedParameterTypes,
|
||||
ContainerUtil.concat(requiredPositionalArgumentTypes,
|
||||
optionalPositionalArgumentTypes,
|
||||
positionalVarargArgumentTypes));
|
||||
}
|
||||
|
||||
public static @Nullable PyTypeParameterMapping mapByShape(@NotNull List<? extends PyType> expectedTypes,
|
||||
@NotNull List<? extends PyType> actualTypes,
|
||||
Option @NotNull ... options) {
|
||||
|
||||
@@ -0,0 +1,602 @@
|
||||
// Copyright 2000-2025 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
|
||||
package com.jetbrains.python
|
||||
|
||||
import com.jetbrains.python.fixtures.PyTestCase
|
||||
import com.jetbrains.python.psi.PyFunction
|
||||
import com.jetbrains.python.psi.PyTargetExpression
|
||||
import com.jetbrains.python.psi.types.PyCallableType
|
||||
import com.jetbrains.python.psi.types.PyTypeChecker
|
||||
import com.jetbrains.python.psi.types.TypeEvalContext
|
||||
|
||||
class PyCallableParameterMappingTest : PyTestCase() {
|
||||
|
||||
fun testEmpty() {
|
||||
checkMatch("()", "()")
|
||||
checkMatch("(/)", "(/)")
|
||||
checkMatch("()", "(/)")
|
||||
checkMatch("(/)", "()")
|
||||
}
|
||||
|
||||
fun testSimple() {
|
||||
checkMatch("(a)", "(a)")
|
||||
checkMatch("(a, b, c)", "(a, b, c)")
|
||||
checkNotMatch("(a: int)", "(b: int)")
|
||||
checkMatch("(a: int)", "(a: int, b: str = 'default')")
|
||||
}
|
||||
|
||||
fun testKeywordOnlyVsStandard() {
|
||||
checkMatch("(*, a: int)", "(a: int)")
|
||||
checkNotMatch("(*, a: int)", "(b: int)") // extra b, missing a (names don't match)
|
||||
checkNotMatch("(*, a: int)", "(a: str)") // names are OK, but types are incompatible
|
||||
checkNotMatch("(*, a: int, b: str)", "(a: int)") // missing parameter b
|
||||
checkNotMatch("(*, a: int)", "(a: int, b: str)") // extra parameter b
|
||||
checkNotMatch("(*, a: int, b: str)", "(a: int, b: str, c: float)") // extra parameter c
|
||||
checkMatch("(*, a: int, b: str)", "(a: int, b: str)") // OK
|
||||
checkMatch("(*, b: str, a: int)", "(a: int, b: str)") // OK, order is not important
|
||||
checkMatch("(*, a: int, b: str)", "(b: str, a: int)") // OK, order is not important
|
||||
checkNotMatch("(*, b: int, a: str)", "(x: str, y: int)") // Names don't match
|
||||
}
|
||||
|
||||
fun testKeywordOnlyVsStandardWithDefaults() {
|
||||
checkMatch("(*, a: int)", "(a: int = 1)") // OK
|
||||
checkMatch("(*, a: int)", "(a: int, b: str = 'str')") // OK, b ignored
|
||||
checkMatch("(*, a: int, b: str)", "(a: int, b: str, c: float = 1.0)") // OK, c ignored
|
||||
checkNotMatch("(*, a: int, b: str = 'str')", "(a: int, b: str)") // b is missing default
|
||||
checkMatch("(*, a: int, b: str = 'str')", "(a: int, b: str = 'str')") // OK, both have defaults
|
||||
checkNotMatch("(*, a: int = 1, b: str = 'str')", "(a: int, b: str = 'str')") // default in actual a is missing
|
||||
checkMatch("(*, a: int = 1, b: str = 'str')", "(a: int = 1, b: str = 'str')") // OK
|
||||
}
|
||||
|
||||
fun testKeywordOnlyVsKeywordOnly() {
|
||||
checkMatch("(*, a: int, b: str)", "(*, b: str, a: int)") // OK, order is not important
|
||||
checkNotMatch("(*, a: int, b: str)", "(*, x: str, y: int)") // Names mismatch
|
||||
checkMatch("(*, a: int, b: str)", "(*, a: int, b: str)") // Same signatures
|
||||
checkMatch("(*, a: int, b: str, c: float)", "(*, c: float, a: int, b: str)") // OK, order is not important
|
||||
checkNotMatch("(*, a: int, b: str, c: float)", "(*, c: float, a: bool, b: str)") // Types mismatch
|
||||
}
|
||||
|
||||
fun testKeywordOnlyVsKeywordOnlyWithDefault() {
|
||||
checkNotMatch("(*, a: int = 1)", "(a: int)") // a is missing default
|
||||
checkMatch("(*, a: int)", "(*, a: int = 1)") // OK
|
||||
checkNotMatch("(*, a: int, b: str = 'default')", "(*, a: int, b: str)") // b is missing default
|
||||
checkMatch("(*, a: int, b: str)", "(*, a: int, b: str = 'default')") // OK
|
||||
checkMatch("(*, a: int = 0, b: str = 'default')", "(*, a: int = 1, b: str = 'other')") // default val is not important
|
||||
checkMatch("(*, a: int, b: str)", "(*, a: int, b: str = 'default', c: float = 1.0)") // c is not important unless it has a default
|
||||
checkNotMatch("(*, a: int, b: str)", "(*, a: int, b: str = 'default', c: float)") // extra param c
|
||||
checkNotMatch("(*, a: int, b: str = 'default', c: float)", "(*, a: int, b: str, c: float)") // b is missing default
|
||||
checkNotMatch("(*, a: int, b: str = 'str', c: bool = True)", "(*, a: int)") // b, c are missing
|
||||
checkMatch("(*, a: int)", "(*, a: int, b: str = 'str', c: bool = True)") // OK, all missing have defaults
|
||||
}
|
||||
|
||||
fun testKeywordOnlyVsPositionalOnly() {
|
||||
checkNotMatch("(*, a: int)", "(a: int, /)") // Missing keyword-only parameter a
|
||||
checkNotMatch("(*, a: int = 1)", "(a: int, /)") // Missing keyword-only parameter a
|
||||
checkNotMatch("(*, a: int = 1)", "(a: int = 1, /)") // Missing keyword-only parameter a
|
||||
}
|
||||
|
||||
fun testStandardVsKeywordOnly() {
|
||||
checkNotMatch("(a: int)", "(*, a: int)") // extra parameter a
|
||||
checkNotMatch("(a: int)", "(*, a: int, b: str)") // extra parameters a, b
|
||||
checkNotMatch("(a: int)", "(*, a: int, b: str, c: float)") // extra parameters a, b, c
|
||||
}
|
||||
|
||||
fun testStandardVsKeywordOnlyWithDefaults() {
|
||||
checkMatch("()", "(x: int = 0)") // OK, x ignored
|
||||
checkNotMatch("(a: int = 1)", "(*, a: int)") // extra parameter a
|
||||
checkNotMatch("(a: int)", "(*, a: int = 1)") // too many positionals
|
||||
checkNotMatch("(a: int = 1)", "(*, a: int = 1)") // too many positionals
|
||||
}
|
||||
|
||||
fun testStandardVsStandard() {
|
||||
checkNotMatch("()", "(a: int)") // a is missing
|
||||
checkMatch("(a: int)", "(a: int)")
|
||||
checkMatch("(a: int, b: str)", "(a: int, b: str)")
|
||||
checkNotMatch("(a: int, b: str)", "(b: int, a: str)") // names don't match (b vs a, a vs b)
|
||||
checkNotMatch("(a: int, b: str)", "(a: str, b: int)") // types don't match
|
||||
checkNotMatch("(a: int, b: str)", "()") // too many positionals
|
||||
checkNotMatch("()", "(a: int, b: str)") // extra parameters a, b
|
||||
}
|
||||
|
||||
fun testStandardVsStandardWithDefaults() {
|
||||
checkMatch("(a: int, b: str)", "(a: int, b: str, c: bool = True)") // OK, c ignored
|
||||
checkNotMatch("(a: int, b: str, c: bool = True)", "(a: int, b: str)") // expected 2 but received 3
|
||||
checkMatch("(a: int = 1, b: str = 'str')", "(a: int = 1, b: str = 'str')") // Ok
|
||||
checkNotMatch("(a: int = 1, b: str = 'str')", "(a: int = 1, b: str)") // default in actual b is missing
|
||||
checkMatch("(a: int, b: str = 'str')", "(a: int = 1, b: str = 'str')") // OK
|
||||
checkMatch("(a: int, b: str)", "(a: int = 1, b: str = 'str')") // OK
|
||||
}
|
||||
|
||||
fun testStandardVsPosOnly() {
|
||||
checkNotMatch("(a: int)", "(a: int, /)") // a is not pos-only
|
||||
checkNotMatch("(a: int, b: str)", "(a: int, /)") // a, b are not pos-only
|
||||
}
|
||||
|
||||
fun testStandardVsPosOnlyWithDefaults() {
|
||||
checkNotMatch("(a: int)", "(a: int = 1, /)")
|
||||
checkNotMatch("(a: int = 1)", "(a: int = 1, /)")
|
||||
checkNotMatch("(a: int = 1)", "(a: int, /)")
|
||||
}
|
||||
|
||||
fun testPosOnlyVsStandard() {
|
||||
checkMatch("(a: int, /)", "(a: int)") // OK
|
||||
checkMatch("(a: int, /)", "(b: int)") // OK, name doesn't matter
|
||||
checkMatch("(a: int, b: str, /)", "(x: int, y: str)") // OK, names don't matter
|
||||
checkNotMatch("(a: int, b: str, /)", "(x: int, y: bool)") // Types don't match
|
||||
checkNotMatch("(a: int, b: str, /)", "(x: int)") // Function accepts too many positional parameters; expected 1 but received 2
|
||||
checkNotMatch("(a: int, /)", "(x: int, y: str)") // Extra param y
|
||||
}
|
||||
|
||||
fun testPosOnlyVsStandardWithDefaults() {
|
||||
checkMatch("(a: int, /)", "(a: int = 1)") // OK
|
||||
checkMatch("(a: int = 1, /)", "(a: int = 1)") // OK
|
||||
checkNotMatch("(a: int = 1, /)", "(a: int)") // a in actual is missing default
|
||||
checkMatch("(a: int, b: str, /)", "(x: int, y: str = 'str')") // OK
|
||||
checkMatch("(a: int, b: str, /)", "(x: int, y: str = 'str')") // OK
|
||||
checkMatch("(a: int, b: str, /)", "(x: int = 1, y: str = 'str')") // OK
|
||||
checkMatch("(a: int, b: str = 'str', /)", "(x: int, y: str = 'str')") // OK
|
||||
checkNotMatch("(a: int = 1, b: str = 'str', /)", "(x: int, y: str = 'str')") // x is missing default
|
||||
}
|
||||
|
||||
fun testPosOnlyVsPosOnly() {
|
||||
checkMatch("(/)", "(/)") // OK
|
||||
checkMatch("(a: int, /)", "(a: int, /)") // OK
|
||||
checkNotMatch("(a: int, /)", "(a: str, /)") // Types mismatch
|
||||
checkMatch("(a: int, b: str, /)", "(x: int, y: str, /)") // Names are not important
|
||||
checkNotMatch("(a: int, b: str, /)", "(x: int, /)") // Too many positionals
|
||||
checkNotMatch("(a: int, b: str, /)", "(x: int, /)") // Too few positionals
|
||||
}
|
||||
|
||||
fun testPosOnlyVsPosOnlyWithDefaults() {
|
||||
checkMatch("(a: int, /)", "(a: int = 1, /)") // OK
|
||||
checkMatch("(a: int = 1, /)", "(a: int = 1, /)") // OK
|
||||
checkMatch("(a: int = 1, /)", "(x: int = 1, /)") // OK, names are not important
|
||||
checkNotMatch("(a: int = 1, /)", "(a: int, /)") // a is missing default
|
||||
checkNotMatch("(a: int = 1, /)", "(x: int, /)") // x is missing default
|
||||
checkNotMatch("(a: int, b: str = 'default', /)", "(a: int, b: str, /)") // b is missing default
|
||||
checkMatch("(a: int, b: str, /)", "(a: int, b: str = 'default', /)") // OK
|
||||
}
|
||||
|
||||
fun testPosOnlyVsKeywordOnly() {
|
||||
checkNotMatch("(a: int, /)", "(*, a: int)")
|
||||
checkNotMatch("(a: int, b: str, /)", "(*, a: int, b: str)")
|
||||
checkNotMatch("(a: int = 1, b: str = 'str', /)", "(*, a: int = 1, b: str = 'str')")
|
||||
}
|
||||
|
||||
fun testUnionType() {
|
||||
checkMatch("(a: int)", "(a: int | str)")
|
||||
checkNotMatch("(a: int | str)", "(a: int)")
|
||||
}
|
||||
|
||||
fun testKeywordOnlyWithKwargs() {
|
||||
checkNotMatch("(*, a: int, b: str)", "(**kwargs: int)") // types mismatch
|
||||
checkMatch("(*, a: int, b: str)", "(**kwargs: int | str)") // OK
|
||||
checkMatch("(*, a: int, b: str)", "(*, a: int, **kwargs: str)") // OK
|
||||
checkNotMatch("(**kwargs: int | str)", "(*, a: int = 1, **kwargs: str)") // types mismatch
|
||||
checkNotMatch("(**kwargs: int | str)", "(**kwargs: int)") // types mismatch
|
||||
checkMatch("(*, a: int, **kwargs: str)", "(**kwargs: int | str)") // OK
|
||||
checkNotMatch("(*, a: int, **kwargs: str)", "(**kwargs: int)") // types mismatch
|
||||
checkMatch("(**kwargs: int)", "(**kwargs: int | str)") // OK
|
||||
checkNotMatch("(**kwargs: int)", "(*, a: int = 1, **kwargs: str)") // types mismatch
|
||||
checkNotMatch("(a: int, b: str)", "(**kwargs: int | str)") // not keyword-only
|
||||
checkNotMatch("(a: int, b: str)", "(*, a: int, **kwargs: str)") // not keyword-only
|
||||
}
|
||||
|
||||
fun testStandardVsVararg() {
|
||||
checkNotMatch("(a: int, *args)", "(a: int)")
|
||||
checkMatch("(a: int)", "(a: int, *args)")
|
||||
checkNotMatch("(a: int, *args)", "(a: int, b: str)")
|
||||
checkNotMatch("(a: int, b: str)", "(a: int, *args)")
|
||||
}
|
||||
|
||||
fun testArgsVsStandardOrPositionalWithArgs() {
|
||||
checkMatch("(*args)", "(a = 1, *args)") // OK
|
||||
checkNotMatch("(*args)", "(a, *args)") // Missing default
|
||||
checkMatch("(*args: int)", "(a: int = 1, *args: int)") // OK
|
||||
checkNotMatch("(*args: int)", "(a: int, *args: int)") // Missing default
|
||||
checkNotMatch("(*args: int)", "(a: int, b: int = 1, *args: int)") // Missing default
|
||||
checkMatch("(*args: int)", "(a: int = 1, b: int = 1, *args: int)") // OK
|
||||
checkMatch("(*args: int)", "(a: int = 1, /, *args: int)") // OK
|
||||
checkNotMatch("(*args: int)", "(a: int, /, *args: int)") // Missing default
|
||||
checkNotMatch("(*args: int)", "(a: int, b: int = 1, /, *args: int)") // Missing default
|
||||
checkMatch("(*args: int)", "(a: int = 1, b: int = 1, /, *args: int)") // OK
|
||||
}
|
||||
|
||||
fun testKwargsVsStandardOrKeywordOnly() {
|
||||
checkMatch("(**kwargs)", "(a = 1, **kwargs)") // OK
|
||||
checkNotMatch("(**kwargs)", "(a, **kwargs)") // Missing default
|
||||
checkMatch("(**kwargs: int)", "(a: int = 1, **kwargs: int)") // OK
|
||||
checkNotMatch("(**kwargs: int)", "(a: int, **kwargs: int)") // Missing default
|
||||
checkNotMatch("(**kwargs: int)", "(a: int, b: int = 1, **kwargs: int)") // Missing default
|
||||
checkMatch("(**kwargs: int)", "(a: int = 1, b: int = 1, **kwargs: int)") // OK
|
||||
checkMatch("(**kwargs: int)", "(*, a: int = 1, **kwargs: int)") // OK
|
||||
checkNotMatch("(**kwargs: int)", "(*, a: int, **kwargs: int)") // Missing default
|
||||
checkNotMatch("(**kwargs: int)", "(*, a: int, b: int = 1, **kwargs: int)") // Missing default
|
||||
checkMatch("(**kwargs: int)", "(*, a: int = 1, b: int = 1, **kwargs: int)") // OK
|
||||
}
|
||||
|
||||
fun testMixedPositionalAndKeyword() {
|
||||
checkMatch("(a: int, b: str, /, c: float, *, d: bool)", "(a: int, b: str, c: float, d: bool)")
|
||||
checkNotMatch("(a: int, b: str, c: float, d: bool)", "(a: int, b: str, /, c: float, *, d: bool)")
|
||||
}
|
||||
|
||||
fun testStandardBetweenPosOnlyAndKwOnly() {
|
||||
checkMatch("(a: int, /, b: str, *args, c: float)", "(a: int, /, b: str, *args, c: float = 1.0)")
|
||||
checkMatch("(a: int, /, b: str, c: str, *args, d: float)", "(a: int, /, b: str, c: str, *args, d: float = 1.0)")
|
||||
checkMatch("(a: int, /, b: str, c: str, *args, d: float)", "(a: int, /, b: str, c: str, c1: str = 'str', *args, d: float = 1.0)")
|
||||
checkNotMatch("(a: int, /, b: str, *args, c: float = 1.0)", "(a: int, /, b: str, *args, c: float)") // c is missing default
|
||||
checkMatch("(a: int, /, b: str = 'str', *args, c: float)", "(a: int, /, b: str = 'str', *args, c: float = 1.0)")
|
||||
}
|
||||
|
||||
fun testComplexSignatures() {
|
||||
// OK
|
||||
checkMatch("(a: int, b: str, /, c: float, *, d: bool)",
|
||||
"(a: int, b: str, c: float, d: bool)")
|
||||
// **kwargs is missing
|
||||
checkNotMatch("(a: int, b: str, /, c: float, *, d: bool, **kwargs)",
|
||||
"(a: int, b: str, c: float, d: bool)")
|
||||
// OK
|
||||
checkMatch("(a: int, b: str, /, c: float, *, d: bool, **kwargs)",
|
||||
"(a: int, b: str, c: float, d: bool, **kwargs)")
|
||||
// missing keyword a, b
|
||||
checkNotMatch("(a: int, b: str, c: float, d: bool)",
|
||||
"(a: int, b: str, /, c: float, *, d: bool)")
|
||||
// extra d
|
||||
checkNotMatch("(a: int, b: str, c: float, d: bool)",
|
||||
"(a: int, b: str, /, c: float, *, d: bool, **kwargs)")
|
||||
// missing keyword d
|
||||
checkNotMatch("(a: int, b: str, /, c: float, *args, d: bool, **kwargs)",
|
||||
"(a: int, b: str, c: float, d: bool)")
|
||||
// OK, d is keyword
|
||||
checkNotMatch("(a: int, b: str, /, c: float, *args, d: bool, **kwargs)",
|
||||
"(a: int, b: str, c: float, *, d: bool)")
|
||||
// OK
|
||||
checkMatch("(a: int, b: str, /, c: float, *args, d: bool, **kwargs)",
|
||||
"(a: int, b: str, c: float, *args: int, d: bool, **kwargs: str)")
|
||||
// type mismatch for d
|
||||
checkNotMatch("(a: int, b: str, /, c: float, *args, d: bool, **kwargs)",
|
||||
"(a: int, b: str, c: float, *args: int, d: str, **kwargs: str)")
|
||||
// keyword-only e: int from expected vs **kwargs: str types mismatch
|
||||
checkNotMatch("(a: int, b: str, /, c: float, *args, d: bool, e: int, **kwargs)",
|
||||
"(a: int, b: str, c: float, *args: int, d: bool, **kwargs: str)")
|
||||
// OK, e matches vs kwargs now
|
||||
checkMatch("(a: int, b: str, /, c: float, *args, d: bool, e: str, **kwargs)",
|
||||
"(a: int, b: str, c: float, *args: int, d: bool, **kwargs: str)")
|
||||
|
||||
// keyword d is missing
|
||||
checkNotMatch("(a: int, b: str = 'default', /, c: float = 1.0, *, d: bool = True)", "(a: int)")
|
||||
// Default for b is missing
|
||||
checkNotMatch("(a: int, /, *, d: bool = True)", "(a: int, b: str = 'default', /, c: float = 1.0, *, d: bool)")
|
||||
// OK
|
||||
checkMatch("(a: int, /, *, d: bool = True)", "(a: int, b: str = 'default', /, c: float = 1.0, *, d: bool = True)")
|
||||
// d: bool vs d: str types mismatch
|
||||
checkNotMatch("(a: int, b: str, /, c: float, *, d: bool)", "(a: int, b: str, c: float, *, d: str)")
|
||||
// OK
|
||||
checkMatch("(a: int, /)", "(a: int, b: str = 'default', /, c: float = 1.0, *, d: bool = True)")
|
||||
}
|
||||
|
||||
fun testKwargsTypeCompatibility() {
|
||||
checkMatch("(*, a: int, b: str)", "(**kwargs: int | str)")
|
||||
checkMatch("(*, a: int, b: str)", "(*, a: int, **kwargs: str)")
|
||||
|
||||
checkNotMatch("(*, a: int, b: str)", "(**kwargs: int)") // str is not a subtype of int
|
||||
checkMatch("(*, a: int, b: int)", "(**kwargs: int)") // int is a subtype of int
|
||||
|
||||
checkMatch("(*, a: int, **kwargs: str)", "(**kwargs: int | str)")
|
||||
checkMatch("(**kwargs: int)", "(**kwargs: int | str)")
|
||||
}
|
||||
|
||||
fun testArgsTypeCompatibility() {
|
||||
checkNotMatch("(*args: int | str)", "(a: int, /, *args: str)")
|
||||
checkMatch("(*args: int)", "(a: int = 1, /, *args: int | str)")
|
||||
checkNotMatch("(*args: int | str)", "(*args: str)")
|
||||
checkNotMatch("(*args: int | str)", "(*args: int)")
|
||||
|
||||
checkMatch("(*args: int)", "(*args: int | str)")
|
||||
checkNotMatch("(*args: int)", "(a: int = 1, /, *args: str)")
|
||||
checkNotMatch("(*args: int)", "(*args: str)")
|
||||
}
|
||||
|
||||
fun testKeyWordOnlyAfterArgs() {
|
||||
checkMatch("(*args, a: int, b: str, c: float)", "(*args, a: int = 1, b: str = 'x', c: float = 1.0)")
|
||||
checkNotMatch("(*args, a: int = 1, b: str = 'x', c: float = 1.0)", "(*args, a: int, b: str, c: float)") // missing defaults
|
||||
checkMatch("(*args, a: int, b: str = 'x', c: float)", "(*args, a: int = 1, b: str = 'y', c: float = 1.0)")
|
||||
checkNotMatch("(*args, a: int, b: str, c: float)", "(*args, a: int, b: str, x: float)") // missing c, extra x
|
||||
}
|
||||
|
||||
fun testIntStrKwargs() {
|
||||
// Test compatibility between **kwargs with different type annotations
|
||||
|
||||
// Union types in **kwargs
|
||||
checkNotMatch("(**kwargs: int | str)", "(**kwargs: int)") // int | str is not a subtype of int
|
||||
|
||||
// **kwargs with keyword-only parameters
|
||||
checkNotMatch("(**kwargs: int | str)", "(*, a: int = 1, **kwargs: str)")
|
||||
checkNotMatch("(**kwargs: int | str)", "(*, a: int = 1, **kwargs: int | str)")
|
||||
checkMatch("(**kwargs: int | str)", "(*, a: int | str = 1, **kwargs: int | str)")
|
||||
|
||||
// Same type in both **kwargs
|
||||
checkMatch("(**kwargs: int | str)", "(**kwargs: int | str)")
|
||||
}
|
||||
|
||||
fun testCallableParameterTypeContravariance() {
|
||||
checkMatch("(a: int)", "(a: float)")
|
||||
checkNotMatch("(a: float)", "(a: int)")
|
||||
|
||||
// *args type mismatch
|
||||
checkNotMatch("(*args: int)", "(a: int = 1, /, *args: str)")
|
||||
checkNotMatch("(*args: int)", "(*args: str)")
|
||||
checkNotMatch("(*args: int | str)", "(a: int = 1, /, *args: str)")
|
||||
checkNotMatch("(*args: int | str)", "(*args: str)")
|
||||
checkNotMatch("(*args: int | str)", "(*args: int)")
|
||||
checkNotMatch("(a: int, /, *args: str)", "(*args: int)")
|
||||
|
||||
// **kwargs type mismatch
|
||||
checkNotMatch("(*, a: int, b: str)", "(**kwargs: int)") // str is not a subtype of int
|
||||
checkNotMatch("(**kwargs: int | str)", "(**kwargs: int)") // int | str is not a subtype of int
|
||||
checkNotMatch("(**kwargs: int | str)", "(*, a: int = 1, **kwargs: str)")
|
||||
checkNotMatch("(**kwargs: int | str)", "(*, a: int = 1, **kwargs: int | str)")
|
||||
checkNotMatch("(**kwargs: float)", "(**kwargs: int)")
|
||||
checkNotMatch("(*, a: int, **kwargs: str)", "(**kwargs: int)")
|
||||
checkNotMatch("(**kwargs: int)", "(*, a: int, **kwargs: str)")
|
||||
|
||||
// *args type mismatch in positional-only context
|
||||
checkNotMatch("(a: int, b: str, /)", "(*args: int)")
|
||||
|
||||
// Complex signature with type mismatch in keyword-only parameter
|
||||
checkNotMatch("(a: int, b: str, /, c: float, *args, d: bool, **kwargs)",
|
||||
"(a: int, b: str, c: float, *args: int, d: str, **kwargs: str)")
|
||||
}
|
||||
|
||||
fun testArgsParameter() {
|
||||
// If a callable B has a signature with a *args parameter, callable A
|
||||
// must also have a *args parameter to be a subtype of B, and the type of
|
||||
// B's *args parameter must be a subtype of A's *args parameter
|
||||
|
||||
checkMatch("()", "(*args: int)")
|
||||
checkMatch("()", "(*args: float)")
|
||||
|
||||
checkMatch("(a: int)", "(a: int, *args: str)")
|
||||
checkNotMatch("(a: int, *args)", "(a: int)")
|
||||
|
||||
checkNotMatch("(*args: int)", "()")
|
||||
checkMatch("(*args: int)", "(*args: float)")
|
||||
|
||||
checkNotMatch("(*args: float)", "()")
|
||||
checkNotMatch("(*args: float)", "(*args: int)")
|
||||
}
|
||||
|
||||
fun testPositionalOnlyWithArgs() {
|
||||
// If a callable B has a signature with one or more positional-only parameters,
|
||||
// a callable A is a subtype of B if A has an *args parameter whose type is a
|
||||
// supertype of the types of any otherwise-unmatched positional-only parameters in B
|
||||
|
||||
checkNotMatch("(a: int, b: str, /)", "(*args: int)")
|
||||
checkMatch("(a: int, b: str, /)", "(*args: int | str)")
|
||||
checkMatch("(a: int, b: str, /)", "(a: int, /, *args: str)")
|
||||
|
||||
checkNotMatch("(*args: int | str)", "(a: int, /, *args: str)")
|
||||
checkNotMatch("(*args: int | str)", "(*args: int)")
|
||||
checkMatch("(a: int, b: str, /, *args: str)", "(*args: int | str)")
|
||||
checkMatch("(a: int, /, *args: str)", "(*args: int | str)")
|
||||
checkNotMatch("(a: int, /, *args: str)", "(*args: int)")
|
||||
checkMatch("(*args: int)", "(*args: int | str)")
|
||||
checkNotMatch("(*args: int)", "(a: int, /, *args: str)")
|
||||
|
||||
checkNotMatch("(a: int, b: str)", "(*args: int | str)")
|
||||
checkNotMatch("(a: int, b: str)", "(a: int, /, *args: str)")
|
||||
}
|
||||
|
||||
fun testKwargsParameter() {
|
||||
checkMatch("()", "(**kwargs: int)")
|
||||
checkMatch("()", "(**kwargs: float)")
|
||||
|
||||
checkNotMatch("(**kwargs: int)", "()")
|
||||
checkMatch("(**kwargs: int)", "(**kwargs: float)")
|
||||
|
||||
checkNotMatch("(**kwargs: float)", "()")
|
||||
checkNotMatch("(**kwargs: float)", "(**kwargs: int)")
|
||||
}
|
||||
|
||||
fun testEmptyWithArgs() {
|
||||
checkMatch("()", "(*args)")
|
||||
checkNotMatch("(*args)", "()")
|
||||
}
|
||||
|
||||
fun testMixedParametersWithDefault() {
|
||||
checkMatch("(a: int, b: str, /)", "(a: int, b: str = 'default', /)") // OK
|
||||
checkNotMatch("(a: int, b: str = 'default', /)", "(a: int, b: str, /)") // b is missing default
|
||||
checkMatch("(a: int, *, b: str)", "(a: int = 0, *, b: str)") // OK
|
||||
checkNotMatch("(a: int = 0, *, b: str)", "(a: int, *, b: str)") // a is missing default
|
||||
checkMatch("(a: int, /, b: str, *, c: float)", "(a: int = 0, /, b: str = 'default', *, c: float = 1.0)") // OK
|
||||
checkNotMatch("(a: int = 0, /, b: str = 'default', *, c: float = 1.0)",
|
||||
"(a: int, /, b: str, *, c: float)") // a, b, c are missing default
|
||||
}
|
||||
|
||||
fun testDefaultWithArgsKwargs() {
|
||||
checkMatch("(a: int, *args)", "(a: int = 0, *args)") // OK
|
||||
checkNotMatch("(a: int = 0, *args)", "(a: int, *args)") // a is missing default
|
||||
checkMatch("(a: int, **kwargs)", "(a: int = 0, **kwargs)") // OK
|
||||
checkNotMatch("(a: int = 0, **kwargs)", "(a: int, **kwargs)") // a is missing default
|
||||
checkMatch("(*, a: int, **kwargs)", "(*, a: int = 0, **kwargs)") // OK
|
||||
checkNotMatch("(*, a: int = 0, **kwargs)", "(*, a: int, **kwargs)") // a is missing default
|
||||
checkMatch("(a: int, /, b: str, *args, c: float, **kwargs)",
|
||||
"(a: int = 0, /, b: str = 'default', *args, c: float = 1.0, **kwargs)") // OK
|
||||
checkNotMatch("(a: int = 0, /, b: str = 'default', *args, c: float = 1.0, **kwargs)",
|
||||
"(a: int, /, b: str, *args, c: float, **kwargs)") // a, b, c are missing default
|
||||
}
|
||||
|
||||
fun testArgsKwargs() {
|
||||
checkMatch("(*args: int, **kwargs: str)", "(*args: int, **kwargs: str)")
|
||||
checkNotMatch("(*args: int, **kwargs: str)", "(*args: int, **kwargs: int)")
|
||||
checkNotMatch("(*args: str, **kwargs: str)", "(*args: int, **kwargs: int)")
|
||||
checkMatch("(a: int, *args: int, **kwargs: str)", "(a: int, *args: int, **kwargs: str)")
|
||||
checkNotMatch("(a: int, *args: str, **kwargs: str)", "(a: str, *args: int, **kwargs: str)")
|
||||
}
|
||||
|
||||
fun testArgsKwargsWithTuple() {
|
||||
checkMatch("(*args: *tuple[int, ...], **kwargs: str)", "(a: int = 1, b: int = 1, *args: *tuple[int, ...], **kwargs: str)")
|
||||
checkMatch("(*args: *tuple[int, str], **kwargs: str)", "(*args: *tuple[int, str], **kwargs: str)")
|
||||
checkNotMatch("(*args: *tuple[int, str], **kwargs: str)", "(*args: *tuple[int, str], **kwargs: int)")
|
||||
checkMatch("(*args: *tuple[int, str], **kwargs: int)", "(a: int = 1, b: str = '', **kwargs: int)")
|
||||
// positional arguments for `a` and `b` map to `*args`, while keyword arguments map to `**kwargs`
|
||||
checkNotMatch("(a: int, b: str, **kwargs: int)", "(*args: *tuple[int, str], **kwargs: int)")
|
||||
checkMatch("(a: int, b: str, /, **kwargs: int)", "(*args: *tuple[int, str], **kwargs: int)")
|
||||
checkMatch("(*args: *tuple[int, ...], **kwargs: str)", "(a: int = 1, b: int = 1, *args: *tuple[int, ...], **kwargs: str)")
|
||||
}
|
||||
|
||||
fun testDefaultValueEdgeCases() {
|
||||
checkMatch("()", "(a: int = 0, b: str = 'default', c: float = 1.0)")
|
||||
checkMatch("(a: int)", "(a: int = 0, b: str = 'default', c: float = 1.0)")
|
||||
checkNotMatch("(a: int = 0, b: str = 'default', c: float = 1.0)", "(a: int, b: str)")
|
||||
checkMatch("(a: int, b: str, c: float)", "(a: int, b: str, c: float = 1.0)")
|
||||
checkNotMatch("(a: int, b: str, c: float = 1.0)", "(a: int, b: str, c: float)")
|
||||
checkMatch("(a: int, b: str, c: float)", "(a: int = 0, b: str, c: float)")
|
||||
checkNotMatch("(a: int = 0, b: str, c: float)", "(a: int, b: str, c: float)")
|
||||
}
|
||||
|
||||
fun testDunderGetParamsConsideredPosOnly() {
|
||||
checkMatch("(__a)", "(b)")
|
||||
checkMatch("(__k, __v)", "(a, b)") // Names are not important here
|
||||
checkNotMatch("(a, __k)", "(a, b)") // Should fail as `(a, __k)` is not a valid signature
|
||||
checkNotMatch("(__a, b, __c)", "(a, b, c)") // Should fail as `(__a, b, __c)` is not a valid signature
|
||||
checkMatch("(__a: int, __b: str)", "(*args: int | str)")
|
||||
}
|
||||
|
||||
fun testCallableWithTypeVarTupleVsRealFunc() {
|
||||
checkMatch("[int, str, *Ts, str]", "(a: int, b: str, c: int, d: int, e: float, x: str, y: str)")
|
||||
checkMatch("[int, str, *Ts, str]", "(a: int, b: str, c: int, d: int, e: float, x: str)")
|
||||
checkMatch("[int, str, *Ts, str]", "(a: int, b: str, c: int, d: int, x: str)")
|
||||
checkMatch("[int, str, *Ts, str]", "(a: int, b: str, c: int, x: str)")
|
||||
checkMatch("[int, str, *Ts, str]", "(a: int, b: str, x: str)")
|
||||
checkNotMatch("[int, str, *Ts, str]", "(a: int, b: str, x: int)")
|
||||
checkNotMatch("[int, str, *Ts, str]", "(a: bool, b: str, x: str)")
|
||||
checkNotMatch("[str, str, *Ts, str]", "(a: int, b: str, x: str)")
|
||||
checkNotMatch("[int, str, *Ts, str]", "(a: int, b: str)")
|
||||
checkNotMatch("[int, str, *Ts, str]", "(a: int)")
|
||||
checkMatch("[int, str, *Ts]", "(a: int, b: str)")
|
||||
checkMatch("[int, str, *Ts]", "(a: int, b: str, /)")
|
||||
checkMatch("[*Ts]", "(a: int)")
|
||||
checkMatch("[*Ts]", "()")
|
||||
checkMatch("[*Ts]", "(a: int, b: str)")
|
||||
checkMatch("[*Ts]", "(a: int, b: str, c: float)")
|
||||
checkMatch("[*Ts]", "(a: int, b: str, c: float, /)")
|
||||
}
|
||||
|
||||
fun testCallableWithTypeVarTupleVsCallableWithTypeVarTuple() {
|
||||
checkMatch("[int, str, *Ts, str]", "[int, str, *Ts, str]")
|
||||
checkMatch("[int, str, *Ts, str]", "[int, str, *Ts, bool, str]")
|
||||
checkMatch("[int, str, *Ts, str]", "[int, str, *Ts, bool, int, str]")
|
||||
checkNotMatch("[int, str, *Ts, str]", "[int, str, *Ts]")
|
||||
checkNotMatch("[int, str, *Ts, str]", "[int, *Ts]")
|
||||
checkNotMatch("[int, str, *Ts, str]", "[*Ts]")
|
||||
}
|
||||
|
||||
fun testRealFuncVsTypingCallable() {
|
||||
checkMatch("()", "[]")
|
||||
checkMatch("(a: int, b: str, /)", "[int, str]")
|
||||
checkNotMatch("(a: int, b: str)", "[int, str]")
|
||||
checkMatch("(a: int, b: str, c: bool, /)", "[int, str, bool]")
|
||||
checkNotMatch("(a: int, b: str, c: bool)", "[int, str, bool]")
|
||||
checkNotMatch("(a: int, b: str, c: bool, /)", "[int, str, bool, float]")
|
||||
checkNotMatch("(a: int, b: str, c: bool = True, /)", "[int, str]")
|
||||
checkNotMatch("(a: int, b: str, c: bool = True, /)", "[int, str, bool]")
|
||||
}
|
||||
|
||||
fun testTypingCallableVsTypingCallable() {
|
||||
checkMatch("[]", "[]")
|
||||
|
||||
checkMatch("[int, str]", "[int, str]")
|
||||
checkMatch("[int, str, bool]", "[int, str, bool]")
|
||||
checkNotMatch("[int, str]", "[int]")
|
||||
checkNotMatch("[int]", "[int, str]")
|
||||
checkNotMatch("[int, str, bool]", "[int, str]")
|
||||
|
||||
// Contravariance
|
||||
checkMatch("[int]", "[float]")
|
||||
checkNotMatch("[float]", "[int]")
|
||||
checkMatch("[int, str]", "[float, str]")
|
||||
checkNotMatch("[float, str]", "[int, str]")
|
||||
|
||||
checkMatch("[int]", "[int | str]")
|
||||
checkNotMatch("[int | str]", "[int]")
|
||||
checkMatch("[int, str]", "[int | float, str]")
|
||||
}
|
||||
|
||||
fun testTypingCallableVsRealFunc() {
|
||||
checkMatch("[]", "()")
|
||||
checkMatch("[]", "(a: int = 0)")
|
||||
checkMatch("[int, str]", "(a: int, b: str, /)")
|
||||
checkMatch("[int, str]", "(a: int, b: str)")
|
||||
checkMatch("[int, str, bool]", "(a: int, b: str, c: bool, /)")
|
||||
checkMatch("[int]", "(a: float, /)")
|
||||
checkNotMatch("[float]", "(a: int, /)")
|
||||
checkMatch("[int, str]", "(a: int, b: str = 'default', /)")
|
||||
checkMatch("[int]", "(a: int = 0, /)")
|
||||
}
|
||||
|
||||
fun testPositionalOrKeywordVsWildcardSignature() {
|
||||
checkMatch("(a)", "(*args, **kwargs)")
|
||||
checkMatch("(a, b)", "(a, *args, **kwargs)")
|
||||
// positional-or-keyword "a" is "split" in an illegal way, it maps to a positional-only parameter when passed as a positional argument
|
||||
// or to the keyword-vararg "**kwargs" when passed with a name
|
||||
checkNotMatch("(a)", "(a, /, *args, **kwargs)")
|
||||
checkNotMatch("(a)", "(b, /, *args, **kwargs)")
|
||||
checkNotMatch("(a, b)", "(b, /, *args, **kwargs)")
|
||||
checkNotMatch("(a, *args)", "(a=1, /, *args, **kwargs)")
|
||||
}
|
||||
|
||||
fun testKeywordOnlyVsAnotherOptionalKeywordOnlyAndKwargs() {
|
||||
checkMatch("(*, a)", "(*, b=None, **kwargs)")
|
||||
}
|
||||
|
||||
fun checkNotMatch(expectedSignature: String, actualSignature: String) {
|
||||
checkMatch(expectedSignature, actualSignature, false)
|
||||
}
|
||||
|
||||
fun checkMatch(expectedSignature: String, actualSignature: String, shouldMatch: Boolean = true) {
|
||||
fun String.isTypingCallable() =
|
||||
if (this.firstOrNull() == '[') true
|
||||
else if (this.firstOrNull() == '(') false
|
||||
else error("Invalid signature: $this")
|
||||
|
||||
val expectedIsTypingCallable = expectedSignature.isTypingCallable()
|
||||
val actualIsTypingCallable = actualSignature.isTypingCallable()
|
||||
|
||||
val expectedSubstitution = if (expectedIsTypingCallable)
|
||||
"my_callable1: Callable[$expectedSignature, None]"
|
||||
else
|
||||
"def my_callable1$expectedSignature: ..."
|
||||
|
||||
val actualSubstitution = if (actualIsTypingCallable)
|
||||
"my_callable2: Callable[$actualSignature, None]"
|
||||
else
|
||||
"def my_callable2$actualSignature: ..."
|
||||
|
||||
val fileText = """
|
||||
from typing import Callable, TypeVarTuple
|
||||
Ts = TypeVarTuple('Ts')
|
||||
$expectedSubstitution
|
||||
$actualSubstitution
|
||||
""".trimIndent()
|
||||
|
||||
myFixture.configureByText(PythonFileType.INSTANCE, fileText)
|
||||
|
||||
val expected = myFixture.findElementByText("my_callable1",
|
||||
if (expectedIsTypingCallable)
|
||||
PyTargetExpression::class.java
|
||||
else PyFunction::class.java)
|
||||
assertNotNull(expected)
|
||||
val actual = myFixture.findElementByText("my_callable2",
|
||||
if (actualIsTypingCallable)
|
||||
PyTargetExpression::class.java
|
||||
else PyFunction::class.java)
|
||||
assertNotNull(actual)
|
||||
val context = TypeEvalContext.codeAnalysis(myFixture.getProject(), myFixture.getFile())
|
||||
val expectedCallable = context.getType(expected) as? PyCallableType
|
||||
assertNotNull(expectedCallable)
|
||||
val actualCallable = context.getType(actual) as? PyCallableType
|
||||
assertNotNull(actualCallable)
|
||||
val match = PyTypeChecker.match(expectedCallable!!, actualCallable!!, context)
|
||||
assertEquals(shouldMatch, match)
|
||||
}
|
||||
}
|
||||
@@ -15,6 +15,7 @@ import org.jetbrains.annotations.NotNull;
|
||||
import org.jetbrains.annotations.Nullable;
|
||||
|
||||
import java.util.LinkedHashSet;
|
||||
import java.util.List;
|
||||
import java.util.Set;
|
||||
import java.util.regex.Pattern;
|
||||
|
||||
@@ -267,6 +268,12 @@ public final class PyTypeParameterMappingTest extends PyTestCase {
|
||||
""");
|
||||
}
|
||||
|
||||
public void testMappingVariadicTypeToEmptyParameters() {
|
||||
doTestParameterListMapping("*Ts", "", """
|
||||
*Ts -> *tuple[]
|
||||
""");
|
||||
}
|
||||
|
||||
public void testMappingVariadicTypeToFewPositionalParameters() {
|
||||
doTestParameterListMapping("int, *Ts", "x: int, y: str, z: bool", """
|
||||
int -> int
|
||||
@@ -431,9 +438,11 @@ public final class PyTypeParameterMappingTest extends PyTestCase {
|
||||
PyFunction actualFunction = myFixture.findElementByText("actual", PyFunction.class);
|
||||
PyFunctionType actualFunctionType = assertInstanceOf(context.getType(actualFunction), PyFunctionType.class);
|
||||
|
||||
List<PyCallableParameter> expectedParameters = ContainerUtil.map(ContainerUtil.subList(expectedTupleType.getElementTypes(), 1),
|
||||
type -> PyCallableParameterImpl.nonPsi(type));
|
||||
|
||||
PyTypeParameterMapping mapping =
|
||||
PyTypeParameterMapping.mapWithParameterList(ContainerUtil.subList(expectedTupleType.getElementTypes(), 1),
|
||||
actualFunctionType.getParameters(context), context);
|
||||
PyCallableParameterMapping.mapCallableParameters(expectedParameters, actualFunctionType.getParameters(context), context);
|
||||
|
||||
assertTypeMapping(expectedMapping, mapping, context);
|
||||
}
|
||||
@@ -468,7 +477,7 @@ public final class PyTypeParameterMappingTest extends PyTestCase {
|
||||
assertTypeMapping(expectedMapping, mapping, context);
|
||||
}
|
||||
|
||||
private static void assertTypeMapping(@NotNull String expectedMappingDump,
|
||||
static void assertTypeMapping(@NotNull String expectedMappingDump,
|
||||
@Nullable PyTypeParameterMapping actualMapping,
|
||||
@NotNull TypeEvalContext context) {
|
||||
if (expectedMappingDump.isEmpty()) {
|
||||
|
||||
@@ -1851,17 +1851,35 @@ public class Py3TypeCheckerInspectionTest extends PyInspectionTestCase {
|
||||
|
||||
Ts = TypeVarTuple('Ts')
|
||||
|
||||
|
||||
def foo(a: int, f: Callable[[*Ts], None], args: Tuple[*Ts]) -> None: ...
|
||||
def bar(a: int, b: str) -> None: ...
|
||||
def baz(a: int, b: str, c: float, d: bool) -> None: ...
|
||||
|
||||
|
||||
foo(1, bar, args=(0, 'foo'))
|
||||
foo(1, baz, args=(0, 'foo', 1.0, False))
|
||||
|
||||
foo(1, bar, <warning descr="Expected type 'tuple[int, str]' (matched generic type 'tuple[*Ts]'), got 'tuple[str, int]' instead">args=('foo', 0)</warning>)
|
||||
foo(1, baz, <warning descr="Expected type 'tuple[int, str, float, bool]' (matched generic type 'tuple[*Ts]'), got 'tuple[str, int, float, bool]' instead">args=('foo', 0, 1.0, False)</warning>)
|
||||
""");
|
||||
}
|
||||
|
||||
// PY-53105 TODO investigate
|
||||
//public void testVariadicGenericArgumentByCallableInFunctionMultipleTypeVars() {
|
||||
// doTestByText("""
|
||||
// from typing import Callable, TypeVarTuple, Tuple, TypeVar
|
||||
//
|
||||
// def foo[T, T1, *Ts](a: T, f: Callable[[T, *Ts, T1], None], args: Tuple[*Ts, T, T1]) -> None: ...
|
||||
// def bar(a: int, b: float, c: str, d: bool) -> None: ...
|
||||
// def baz(a: str, b: float, d: int) -> None: ...
|
||||
//
|
||||
// foo(1, bar, args=(1.0, "str", 1, True)) # T -> int, T1 -> bool, *Ts -> (float, str)
|
||||
// foo("str", baz, args=(1.0, "str", 3)) # T - > str, T1 -> int, *Ts -> float
|
||||
// foo(1, <warning descr="Expected type '(int, *Ts, T1) -> None' (matched generic type '(T, *Ts, T1) -> None'), got '(a: str, b: float, d: int) -> None' instead">baz</warning>, <warning descr="Expected type 'tuple[float, int, T1]' (matched generic type 'tuple[*Ts, T, T1]'), got 'tuple[float, str, int]' instead">args=(1.0, "str", 3)</warning>)
|
||||
// foo(1, bar, <warning descr="Expected type 'tuple[float, str, int, bool]' (matched generic type 'tuple[*Ts, T, T1]'), got 'tuple[float, str, float, bool]' instead">args=(1.0, "str", 1.0, True)</warning>)
|
||||
// """);
|
||||
//}
|
||||
|
||||
// PY-53105
|
||||
public void testVariadicGenericCheckCallableInFunction() {
|
||||
doTestByText("""
|
||||
@@ -3818,5 +3836,294 @@ public class Py3TypeCheckerInspectionTest extends PyInspectionTestCase {
|
||||
""");
|
||||
});
|
||||
}
|
||||
|
||||
// Test for callable subtyping rules - covariance and contravariance
|
||||
public void testCallableSubtypingCovarianceContravariance() {
|
||||
doTestByText("""
|
||||
from typing import Callable
|
||||
|
||||
# Test covariance with respect to return types and contravariance with respect to parameter types
|
||||
def func1(
|
||||
cb1: Callable[[float], int],
|
||||
cb2: Callable[[float], float],
|
||||
cb3: Callable[[int], int],
|
||||
) -> None:
|
||||
f1: Callable[[int], float] = cb1 # OK
|
||||
f2: Callable[[int], float] = cb2 # OK
|
||||
f3: Callable[[int], float] = cb3 # OK
|
||||
|
||||
f4: Callable[[float], float] = cb1 # OK
|
||||
f5: Callable[[float], float] = cb2 # OK
|
||||
f6: Callable[[float], float] = <warning descr="Expected type '(float) -> float', got '(int) -> int' instead">cb3</warning> # Error
|
||||
|
||||
f7: Callable[[int], int] = cb1 # OK
|
||||
f8: Callable[[int], int] = <warning descr="Expected type '(int) -> int', got '(float) -> float' instead">cb2</warning> # Error
|
||||
f9: Callable[[int], int] = cb3 # OK
|
||||
""");
|
||||
}
|
||||
|
||||
// https://typing.python.org/en/latest/spec/callables.html#parameter-kinds
|
||||
public void testCallableSubtypingParameterKinds() {
|
||||
doTestByText("""
|
||||
from typing import Protocol
|
||||
|
||||
# Test positional-only, keyword-only, and standard parameters
|
||||
class PosOnly(Protocol):
|
||||
def __call__(self, a: int, b: str, /) -> None: ...
|
||||
|
||||
class KwOnly(Protocol):
|
||||
def __call__(self, *, a: int, b: str) -> None: ...
|
||||
|
||||
class Standard(Protocol):
|
||||
def __call__(self, a: int, b: str) -> None: ...
|
||||
|
||||
def func2(standard: Standard, pos_only: PosOnly, kw_only: KwOnly):
|
||||
f1: Standard = <warning descr="Expected type 'Standard', got 'PosOnly' instead">pos_only</warning> # Error
|
||||
f2: Standard = <warning descr="Expected type 'Standard', got 'KwOnly' instead">kw_only</warning> # Error
|
||||
|
||||
f3: PosOnly = standard # OK
|
||||
f4: PosOnly = <warning descr="Expected type 'PosOnly', got 'KwOnly' instead">kw_only</warning> # Error
|
||||
|
||||
f5: KwOnly = standard # OK
|
||||
f6: KwOnly = <warning descr="Expected type 'KwOnly', got 'PosOnly' instead">pos_only</warning> # Error
|
||||
""");
|
||||
}
|
||||
|
||||
// https://typing.python.org/en/latest/spec/callables.html#args-parameters
|
||||
public void testCallableSubtypingArgsParameter() {
|
||||
doTestByText("""
|
||||
from typing import Protocol
|
||||
|
||||
# Test *args parameter
|
||||
class NoArgs(Protocol):
|
||||
def __call__(self) -> None: ...
|
||||
|
||||
class IntArgs(Protocol):
|
||||
def __call__(self, *args: int) -> None: ...
|
||||
|
||||
class FloatArgs(Protocol):
|
||||
def __call__(self, *args: float) -> None: ...
|
||||
|
||||
def func3(no_args: NoArgs, int_args: IntArgs, float_args: FloatArgs):
|
||||
f1: NoArgs = int_args # OK
|
||||
f2: NoArgs = float_args # OK
|
||||
|
||||
f3: IntArgs = <warning descr="Expected type 'IntArgs', got 'NoArgs' instead">no_args</warning> # Error: missing *args
|
||||
f4: IntArgs = float_args # OK
|
||||
|
||||
f5: FloatArgs = <warning descr="Expected type 'FloatArgs', got 'NoArgs' instead">no_args</warning> # Error: missing *args
|
||||
f6: FloatArgs = <warning descr="Expected type 'FloatArgs', got 'IntArgs' instead">int_args</warning> # Error: float is not subtype of int
|
||||
""");
|
||||
}
|
||||
|
||||
// https://typing.python.org/en/latest/spec/callables.html#args-parameters
|
||||
public void testCallableSubtypingArgsParameter2() {
|
||||
doTestByText("""
|
||||
from typing import Protocol
|
||||
|
||||
class PosOnly(Protocol):
|
||||
def __call__(self, a: int, b: str, /) -> None: ...
|
||||
|
||||
class IntArgs(Protocol):
|
||||
def __call__(self, *args: int) -> None: ...
|
||||
|
||||
class IntStrArgs(Protocol):
|
||||
def __call__(self, *args: int | str) -> None: ...
|
||||
|
||||
class StrArgs(Protocol):
|
||||
def __call__(self, a: int, /, *args: str) -> None: ...
|
||||
|
||||
class Standard(Protocol):
|
||||
def __call__(self, a: int, b: str) -> None: ...
|
||||
|
||||
def func(int_args: IntArgs, int_str_args: IntStrArgs, str_args: StrArgs):
|
||||
f1: PosOnly = <warning descr="Expected type 'PosOnly', got 'IntArgs' instead">int_args</warning> # Error: str is not assignable to int
|
||||
f2: PosOnly = int_str_args # OK
|
||||
f3: PosOnly = str_args # OK
|
||||
f4: IntStrArgs = <warning descr="Expected type 'IntStrArgs', got 'StrArgs' instead">str_args</warning> # Error: int | str is not assignable to str
|
||||
f5: IntStrArgs = <warning descr="Expected type 'IntStrArgs', got 'IntArgs' instead">int_args</warning> # Error: int | str is not assignable to int
|
||||
f6: StrArgs = int_str_args # OK
|
||||
f7: StrArgs = <warning descr="Expected type 'StrArgs', got 'IntArgs' instead">int_args</warning> # Error: str is not assignable to int
|
||||
f8: IntArgs = int_str_args # OK
|
||||
f9: IntArgs = <warning descr="Expected type 'IntArgs', got 'StrArgs' instead">str_args</warning> # Error: int is not assignable to str
|
||||
f10: Standard = <warning descr="Expected type 'Standard', got 'IntStrArgs' instead">int_str_args</warning> # Error: keyword parameters a and b missing
|
||||
f11: Standard = <warning descr="Expected type 'Standard', got 'StrArgs' instead">str_args</warning> # Error: keyword parameter b missing
|
||||
""");
|
||||
}
|
||||
|
||||
// https://typing.python.org/en/latest/spec/callables.html#kwargs-parameters
|
||||
public void testCallableSubtypingKwargsParameters() {
|
||||
doTestByText("""
|
||||
from typing import Protocol
|
||||
|
||||
# Test **kwargs parameter
|
||||
class NoKwargs(Protocol):
|
||||
def __call__(self) -> None: ...
|
||||
|
||||
class IntKwargs(Protocol):
|
||||
def __call__(self, **kwargs: int) -> None: ...
|
||||
|
||||
class FloatKwargs(Protocol):
|
||||
def __call__(self, **kwargs: float) -> None: ...
|
||||
|
||||
def func5(no_kwargs: NoKwargs, int_kwargs: IntKwargs, float_kwargs: FloatKwargs):
|
||||
f1: NoKwargs = int_kwargs # OK
|
||||
f2: NoKwargs = float_kwargs # OK
|
||||
|
||||
f3: IntKwargs = <warning descr="Expected type 'IntKwargs', got 'NoKwargs' instead">no_kwargs</warning> # Error: missing **kwargs
|
||||
f4: IntKwargs = float_kwargs # OK
|
||||
|
||||
f5: FloatKwargs = <warning descr="Expected type 'FloatKwargs', got 'NoKwargs' instead">no_kwargs</warning> # Error: missing **kwargs
|
||||
f6: FloatKwargs = <warning descr="Expected type 'FloatKwargs', got 'IntKwargs' instead">int_kwargs</warning> # Error: float is not subtype of int
|
||||
""");
|
||||
}
|
||||
|
||||
// https://typing.python.org/en/latest/spec/callables.html#kwargs-parameters
|
||||
public void testCallableSubtypingKwargsParameters2() {
|
||||
doTestByText("""
|
||||
from typing import Protocol
|
||||
|
||||
class KwOnly(Protocol):
|
||||
def __call__(self, *, a: int, b: str) -> None: ...
|
||||
|
||||
class IntKwargs(Protocol):
|
||||
def __call__(self, **kwargs: int) -> None: ...
|
||||
|
||||
class IntStrKwargs(Protocol):
|
||||
def __call__(self, **kwargs: int | str) -> None: ...
|
||||
|
||||
class StrKwargs(Protocol):
|
||||
def __call__(self, *, a: int, **kwargs: str) -> None: ...
|
||||
|
||||
class Standard(Protocol):
|
||||
def __call__(self, a: int, b: str) -> None: ...
|
||||
|
||||
def func(int_kwargs: IntKwargs, int_str_kwargs: IntStrKwargs, str_kwargs: StrKwargs):
|
||||
f1: KwOnly = <warning descr="Expected type 'KwOnly', got 'IntKwargs' instead">int_kwargs</warning> # Error: str is not assignable to int
|
||||
f2: KwOnly = int_str_kwargs # OK
|
||||
f3: KwOnly = str_kwargs # OK
|
||||
f4: IntStrKwargs = <warning descr="Expected type 'IntStrKwargs', got 'StrKwargs' instead">str_kwargs</warning> # Error: int | str is not assignable to str
|
||||
f5: IntStrKwargs = <warning descr="Expected type 'IntStrKwargs', got 'IntKwargs' instead">int_kwargs</warning> # Error: int | str is not assignable to int
|
||||
f6: StrKwargs = int_str_kwargs # OK
|
||||
f7: StrKwargs = <warning descr="Expected type 'StrKwargs', got 'IntKwargs' instead">int_kwargs</warning> # Error: str is not assignable to int
|
||||
f8: IntKwargs = int_str_kwargs # OK
|
||||
f9: IntKwargs = <warning descr="Expected type 'IntKwargs', got 'StrKwargs' instead">str_kwargs</warning> # Error: int is not assignable to str
|
||||
f10: Standard = <warning descr="Expected type 'Standard', got 'IntStrKwargs' instead">int_str_kwargs</warning> # Error: Does not accept positional arguments
|
||||
f11: Standard = <warning descr="Expected type 'Standard', got 'StrKwargs' instead">str_kwargs</warning> # Error: Does not accept positional arguments
|
||||
""");
|
||||
}
|
||||
|
||||
// https://typing.python.org/en/latest/spec/callables.html#id4
|
||||
public void testCallableSubtypingDefaultArguments() {
|
||||
doTestByText("""
|
||||
from typing import Protocol
|
||||
|
||||
# Test default arguments
|
||||
class DefaultArg(Protocol):
|
||||
def __call__(self, x: int = 0) -> None: ...
|
||||
|
||||
class NoDefaultArg(Protocol):
|
||||
def __call__(self, x: int) -> None: ...
|
||||
|
||||
class NoX(Protocol):
|
||||
def __call__(self) -> None: ...
|
||||
|
||||
def func8(default_arg: DefaultArg, no_default_arg: NoDefaultArg, no_x: NoX):
|
||||
f1: DefaultArg = <warning descr="Expected type 'DefaultArg', got 'NoDefaultArg' instead">no_default_arg</warning> # Error
|
||||
f2: DefaultArg = <warning descr="Expected type 'DefaultArg', got 'NoX' instead">no_x</warning> # Error
|
||||
|
||||
f3: NoDefaultArg = default_arg # OK
|
||||
f4: NoDefaultArg = <warning descr="Expected type 'NoDefaultArg', got 'NoX' instead">no_x</warning> # Error
|
||||
|
||||
f5: NoX = default_arg # OK
|
||||
f6: NoX = <warning descr="Expected type 'NoX', got 'NoDefaultArg' instead">no_default_arg</warning> # Error
|
||||
""");
|
||||
}
|
||||
|
||||
// https://typing.python.org/en/latest/spec/callables.html#overloads
|
||||
public void testCallableSubtypingOverloads() {
|
||||
doTestByText("""
|
||||
from typing import Protocol
|
||||
|
||||
class Overloaded(Protocol):
|
||||
@overload
|
||||
def __call__(self, x: int) -> int: ...
|
||||
@overload
|
||||
def __call__(self, x: str) -> str: ...
|
||||
|
||||
class IntArg(Protocol):
|
||||
def __call__(self, x: int) -> int: ...
|
||||
|
||||
class StrArg(Protocol):
|
||||
def __call__(self, x: str) -> str: ...
|
||||
|
||||
class FloatArg(Protocol):
|
||||
def __call__(self, x: float) -> float: ...
|
||||
|
||||
def func(overloaded: Overloaded):
|
||||
f1: IntArg = overloaded # OK
|
||||
f2: StrArg = overloaded # OK
|
||||
f3: FloatArg = <warning descr="Expected type 'FloatArg', got 'Overloaded' instead">overloaded</warning> # Error
|
||||
""");
|
||||
}
|
||||
|
||||
// https://typing.python.org/en/latest/spec/callables.html#overloads
|
||||
public void testCallableSubtypingOverloads2() {
|
||||
doTestByText("""
|
||||
from typing import Protocol
|
||||
|
||||
class Overloaded(Protocol):
|
||||
@overload
|
||||
def __call__(self, x: int, y: str) -> float: ...
|
||||
@overload
|
||||
def __call__(self, x: str) -> complex: ...
|
||||
|
||||
class StrArg(Protocol):
|
||||
def __call__(self, x: str) -> complex: ...
|
||||
|
||||
class IntStrArg(Protocol):
|
||||
def __call__(self, x: int | str, y: str = "") -> int: ...
|
||||
|
||||
def func(int_str_arg: IntStrArg, str_arg: StrArg):
|
||||
f1: Overloaded = int_str_arg # OK
|
||||
f2: Overloaded = <warning descr="Expected type 'Overloaded', got 'StrArg' instead">str_arg</warning> # Error
|
||||
""");
|
||||
}
|
||||
|
||||
|
||||
// https://typing.python.org/en/latest/spec/callables.html#signatures-with-paramspecs
|
||||
public void testSignaturesWithParamSpec() {
|
||||
doTestByText("""
|
||||
from typing import Protocol
|
||||
|
||||
class ProtocolWithP[**P](Protocol):
|
||||
def __call__(self, *args: P.args, **kwargs: P.kwargs) -> None: ...
|
||||
|
||||
type TypeAliasWithP[**P] = Callable[P, None]
|
||||
|
||||
def func[**P](proto: ProtocolWithP[P], ta: TypeAliasWithP[P]):
|
||||
# These two types are equivalent
|
||||
f1: TypeAliasWithP[P] = proto # OK
|
||||
f2: ProtocolWithP[P] = ta # OK
|
||||
""");
|
||||
}
|
||||
|
||||
// PY-76883
|
||||
public void testCallableSubtypingKeywordOnlyOrder() {
|
||||
doTestByText("""
|
||||
from typing import Protocol
|
||||
|
||||
class C1(Protocol):
|
||||
def __call__(self, *, a: int, b: str, c: float): ...
|
||||
|
||||
class C2(Protocol):
|
||||
def __call__(self, *, c: float, b: str, a: int): ...
|
||||
|
||||
# Order is not important
|
||||
def foo(c1: C1, c2: C2):
|
||||
_: C1 = c2
|
||||
_: C2 = c1
|
||||
""");
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user