diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyCallableParameterMapping.kt b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyCallableParameterMapping.kt new file mode 100644 index 000000000000..d1692a49e54b --- /dev/null +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyCallableParameterMapping.kt @@ -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 Specification + */ +@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, + actualCallableParameters: List, + 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() + val actualTypes = mutableListOf() + + 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, context: TypeEvalContext): List { + val flattenedParameters = mutableListOf() + 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, context: TypeEvalContext): List? { + val parameters = mutableListOf() + + 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) + } +} \ No newline at end of file diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeChecker.java b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeChecker.java index 09a60b0b8f50..17db3c1f6687 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeChecker.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeChecker.java @@ -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 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; } diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeParameterMapping.java b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeParameterMapping.java index 6dd665a10027..c79a62ad9ec6 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeParameterMapping.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyTypeParameterMapping.java @@ -31,82 +31,6 @@ public final class PyTypeParameterMapping { myMappedTypes = mapping; } - public static @Nullable PyTypeParameterMapping mapWithParameterList(@NotNull List expectedParameterTypes, - @NotNull List actualParameters, - @NotNull TypeEvalContext context) { - List flattenedExpectedParameterTypes = flattenUnpackedTupleTypes(expectedParameterTypes); - int expectedArity = ContainerUtil.exists(flattenedExpectedParameterTypes, Conditions.instanceOf(PyPositionalVariadicType.class)) - ? -1 - : flattenedExpectedParameterTypes.size(); - - List requiredPositionalArgumentTypes = new ArrayList<>(); - List optionalPositionalArgumentTypes = new ArrayList<>(); - List 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 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 expectedTypes, @NotNull List actualTypes, Option @NotNull ... options) { diff --git a/python/testSrc/com/jetbrains/python/PyCallableParameterMappingTest.kt b/python/testSrc/com/jetbrains/python/PyCallableParameterMappingTest.kt new file mode 100644 index 000000000000..4c8a4d4b5c41 --- /dev/null +++ b/python/testSrc/com/jetbrains/python/PyCallableParameterMappingTest.kt @@ -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) + } +} diff --git a/python/testSrc/com/jetbrains/python/PyTypeParameterMappingTest.java b/python/testSrc/com/jetbrains/python/PyTypeParameterMappingTest.java index 68756a5769f1..dc11315f6818 100644 --- a/python/testSrc/com/jetbrains/python/PyTypeParameterMappingTest.java +++ b/python/testSrc/com/jetbrains/python/PyTypeParameterMappingTest.java @@ -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 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()) { diff --git a/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java index 978c16bd949a..83f33d822696 100644 --- a/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java +++ b/python/testSrc/com/jetbrains/python/inspections/Py3TypeCheckerInspectionTest.java @@ -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, args=('foo', 0)) + foo(1, baz, args=('foo', 0, 1.0, False)) """); } + // 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, baz, args=(1.0, "str", 3)) + // foo(1, bar, args=(1.0, "str", 1.0, True)) + // """); + //} + // 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] = cb3 # Error + + f7: Callable[[int], int] = cb1 # OK + f8: Callable[[int], int] = cb2 # 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 = pos_only # Error + f2: Standard = kw_only # Error + + f3: PosOnly = standard # OK + f4: PosOnly = kw_only # Error + + f5: KwOnly = standard # OK + f6: KwOnly = pos_only # 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 = no_args # Error: missing *args + f4: IntArgs = float_args # OK + + f5: FloatArgs = no_args # Error: missing *args + f6: FloatArgs = int_args # 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 = int_args # Error: str is not assignable to int + f2: PosOnly = int_str_args # OK + f3: PosOnly = str_args # OK + f4: IntStrArgs = str_args # Error: int | str is not assignable to str + f5: IntStrArgs = int_args # Error: int | str is not assignable to int + f6: StrArgs = int_str_args # OK + f7: StrArgs = int_args # Error: str is not assignable to int + f8: IntArgs = int_str_args # OK + f9: IntArgs = str_args # Error: int is not assignable to str + f10: Standard = int_str_args # Error: keyword parameters a and b missing + f11: Standard = str_args # 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 = no_kwargs # Error: missing **kwargs + f4: IntKwargs = float_kwargs # OK + + f5: FloatKwargs = no_kwargs # Error: missing **kwargs + f6: FloatKwargs = int_kwargs # 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 = int_kwargs # Error: str is not assignable to int + f2: KwOnly = int_str_kwargs # OK + f3: KwOnly = str_kwargs # OK + f4: IntStrKwargs = str_kwargs # Error: int | str is not assignable to str + f5: IntStrKwargs = int_kwargs # Error: int | str is not assignable to int + f6: StrKwargs = int_str_kwargs # OK + f7: StrKwargs = int_kwargs # Error: str is not assignable to int + f8: IntKwargs = int_str_kwargs # OK + f9: IntKwargs = str_kwargs # Error: int is not assignable to str + f10: Standard = int_str_kwargs # Error: Does not accept positional arguments + f11: Standard = str_kwargs # 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 = no_default_arg # Error + f2: DefaultArg = no_x # Error + + f3: NoDefaultArg = default_arg # OK + f4: NoDefaultArg = no_x # Error + + f5: NoX = default_arg # OK + f6: NoX = no_default_arg # 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 = overloaded # 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 = str_arg # 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 + """); + } }