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:
Daniil Kalinin
2026-01-07 16:01:39 +00:00
committed by intellij-monorepo-bot
co-authored by Mikhail Golubev Petr Golubev
parent 6c42ce2770
commit e2600feb97
6 changed files with 1274 additions and 110 deletions
@@ -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
""");
}
}