From 14e295879d0364032967882ba92cd1c4cf76f20c Mon Sep 17 00:00:00 2001 From: Marcus Mews Date: Mon, 8 Dec 2025 09:55:20 +0000 Subject: [PATCH] PY-85390 Introduce PyExpectedTypeJudgement that aggregates related code - introduces PyExpectedTypeJudgement as a single entry to compute expected types - changes a few clients to use PyExpectedTypeJudgement GitOrigin-RevId: 27734b99b37a364b540f8af985908b282c50afd7 --- .../PyFStringLikeCompletionContributor.java | 11 +- .../PyLiteralTypeCompletionContributor.kt | 36 +- .../typing/PyTypedDictTypeProvider.kt | 2 - .../psi/impl/PyLambdaExpressionImpl.java | 7 +- .../psi/types/PyExpectedTypeJudgement.kt | 459 +++++++ .../python/psi/types/PyTypeChecker.java | 65 +- .../python/PyExpectedTypeJudgmentTest.kt | 1058 +++++++++++++++++ 7 files changed, 1536 insertions(+), 102 deletions(-) create mode 100644 python/python-psi-impl/src/com/jetbrains/python/psi/types/PyExpectedTypeJudgement.kt create mode 100644 python/testSrc/com/jetbrains/python/PyExpectedTypeJudgmentTest.kt diff --git a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/completion/PyFStringLikeCompletionContributor.java b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/completion/PyFStringLikeCompletionContributor.java index 74c2f87de7b2..51ff0d459c14 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/completion/PyFStringLikeCompletionContributor.java +++ b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/completion/PyFStringLikeCompletionContributor.java @@ -5,19 +5,18 @@ import com.intellij.codeInsight.lookup.LookupElement; import com.intellij.codeInsight.lookup.LookupElementDecorator; import com.intellij.openapi.editor.Document; import com.intellij.openapi.project.DumbAware; -import com.intellij.openapi.util.Pair; import com.intellij.openapi.util.text.StringUtil; import com.intellij.patterns.PsiElementPattern; -import com.intellij.psi.PsiElement; import com.intellij.psi.PsiReference; -import com.intellij.psi.util.PsiTreeUtil; import com.intellij.util.ProcessingContext; import com.intellij.util.containers.ContainerUtil; import com.intellij.util.text.CharArrayUtil; import com.jetbrains.python.PyNames; import com.jetbrains.python.psi.*; -import com.jetbrains.python.psi.resolve.PyResolveContext; -import com.jetbrains.python.psi.types.*; +import com.jetbrains.python.psi.types.PyClassType; +import com.jetbrains.python.psi.types.PyExpectedTypeJudgement; +import com.jetbrains.python.psi.types.PyType; +import com.jetbrains.python.psi.types.TypeEvalContext; import org.jetbrains.annotations.NotNull; import java.util.List; @@ -131,7 +130,7 @@ public final class PyFStringLikeCompletionContributor extends CompletionContribu return false; } PyClassType templateType = psiFacade.createClassType(templateClass, false); - PyType expectedType = PyTypeChecker.getExpectedType(stringLiteral, typeEvalContext); + PyType expectedType = PyExpectedTypeJudgement.getExpectedType(stringLiteral, typeEvalContext); return templateType.equals(expectedType); } } diff --git a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/completion/PyLiteralTypeCompletionContributor.kt b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/completion/PyLiteralTypeCompletionContributor.kt index 0556134416ba..d6329395d830 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/completion/PyLiteralTypeCompletionContributor.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/completion/PyLiteralTypeCompletionContributor.kt @@ -5,18 +5,14 @@ import com.intellij.codeInsight.completion.ml.MLRankingIgnorable import com.intellij.codeInsight.lookup.LookupElementBuilder import com.intellij.patterns.PlatformPatterns.psiElement import com.intellij.psi.PsiElement -import com.intellij.psi.util.PsiTreeUtil import com.intellij.ui.IconManager import com.intellij.ui.PlatformIcons import com.intellij.util.ProcessingContext -import com.jetbrains.python.psi.* -import com.jetbrains.python.psi.impl.PyPsiUtils -import com.jetbrains.python.psi.impl.getMappedParameters -import com.jetbrains.python.psi.resolve.PyResolveContext -import com.jetbrains.python.psi.types.PyLiteralType -import com.jetbrains.python.psi.types.PyType -import com.jetbrains.python.psi.types.PyTypeUtil -import com.jetbrains.python.psi.types.TypeEvalContext +import com.jetbrains.python.psi.PyExpression +import com.jetbrains.python.psi.PyReferenceExpression +import com.jetbrains.python.psi.PyStringLiteralExpression +import com.jetbrains.python.psi.StringLiteralExpression +import com.jetbrains.python.psi.types.* /** * Provides literal type variants in the following cases: @@ -46,26 +42,10 @@ private class PyLiteralTypeCompletionProvider : CompletionProvider, result: CompletionResultSet) { diff --git a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypedDictTypeProvider.kt b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypedDictTypeProvider.kt index e32a878422c7..dd810712b81d 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypedDictTypeProvider.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypedDictTypeProvider.kt @@ -13,12 +13,10 @@ import com.jetbrains.python.psi.impl.PyCallExpressionNavigator import com.jetbrains.python.psi.impl.PyEvaluator import com.jetbrains.python.psi.impl.StubAwareComputation import com.jetbrains.python.psi.impl.stubs.PyTypedDictStubImpl -import com.jetbrains.python.psi.resolve.PyResolveContext import com.jetbrains.python.psi.stubs.PyTypedDictFieldStub import com.jetbrains.python.psi.stubs.PyTypedDictStub import com.jetbrains.python.psi.types.* import com.jetbrains.python.psi.types.PyTypedDictType.Companion.TYPED_DICT_TOTAL_PARAMETER -import java.util.* import java.util.stream.Collectors typealias TDFields = LinkedHashMap diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyLambdaExpressionImpl.java b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyLambdaExpressionImpl.java index a2f7fb76924c..a44b563e003f 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyLambdaExpressionImpl.java +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/impl/PyLambdaExpressionImpl.java @@ -10,7 +10,10 @@ import com.jetbrains.python.psi.types.*; import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; -import java.util.*; +import java.util.ArrayList; +import java.util.List; +import java.util.Map; +import java.util.Optional; import static com.intellij.util.containers.ContainerUtil.map; @@ -34,7 +37,7 @@ public class PyLambdaExpressionImpl extends PyElementImpl implements PyLambdaExp } } - @Nullable PyType expected = PyTypeChecker.getExpectedType(this, context); + @Nullable PyType expected = PyExpectedTypeJudgement.getExpectedType(this, context); if (expected instanceof PyCallableType expectedCallable) { var params = new ArrayList(); diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyExpectedTypeJudgement.kt b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyExpectedTypeJudgement.kt new file mode 100644 index 000000000000..acdb52696cbd --- /dev/null +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyExpectedTypeJudgement.kt @@ -0,0 +1,459 @@ +package com.jetbrains.python.psi.types + +import com.intellij.psi.PsiElement +import com.intellij.psi.util.parentOfType +import com.jetbrains.python.PyNames +import com.jetbrains.python.ast.impl.PyPsiUtilsCore.flattenParens +import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider +import com.jetbrains.python.psi.* +import com.jetbrains.python.psi.impl.PyBuiltinCache +import com.jetbrains.python.psi.impl.mapArguments +import com.jetbrains.python.psi.resolve.PyResolveContext +import com.jetbrains.python.psi.types.PyTypeChecker.* + + +object PyExpectedTypeJudgement { + + /** + * Computes the expected type of the given expression from its usage in the AST. + * The expected type is either the explicitly declared type (i.e., type annotation) that constraints the given expression. + * Or the expected type is implied by the grammar and language semantics, e.g., for varargs. + * Note that in some cases the expected type has a circular dependency to itself via type inference, which in turn calls the + * expected type judgment again. This can happen e.g., for cases related to generic functions. + * + * Supported root AST elements: + * - Argument to a call (positional or keyword) + * - Argument to indexed access (i.e., subscription expression) + * - RHS of an assignment + * - Yield expression value (RHS) + * - Return statement value + * + * When given an AST element that is a child C of a root AST element, this method traverses the parent chain upwards + * to compute the expected type of the root element. Based on that, it tries to conclude the expected type of C. + * + * Returns null iff no type declaration, or `Any` was found. + * + * Note: At the moment, this method does not resolve subtype relationships in cases where the AST expectations differ + * from the actual type found. E.g., consider the example `my_var : MyListOfInts = [1]` and suppose that the expected + * type of `1` is requested. This implementation does not (yet) resolve `MyListOfInts` to compute its relation to the + * implied supertype `List`, to then retrieve the type of `List`s type parameter. + */ + @JvmStatic + fun getExpectedType(expr: PyExpression, ctx: TypeEvalContext): PyType? { + // Traverse the AST upwards to find the root expression (either assignment, function call, return statement) + // Do this recursively to easily map the result type to the original sub-element. + // Example: x2: str; x1, (x2, x3) = (42, (expr, "spam")) # expr is the requested sub-element, the whole tuple is the root expression + val parent = expr.parent + when (parent) { + is PyStarArgument, + is PyParenthesizedExpression, + -> { + return getExpectedType(parent, ctx) + } + + is PyAssignmentExpression -> { + val expectedType = fromWalrus(expr) + if (expectedType != null) return expectedType + return getExpectedType(parent, ctx) + } + + is PySliceItem -> { + val cache = PyBuiltinCache.getInstance(expr) + return PyUnionType.union(cache.intType, cache.noneType) + } + + is PyStarExpression -> { + if (parent.parent is PyExpression) { + val typeOfStarParent = getExpectedType(parent.parent as PyExpression, ctx) + if (PyNames.ITERABLE == typeOfStarParent?.name) { + return typeOfStarParent + } + if (typeOfStarParent is PyCollectionType) { + // upcast to Iterable + return createIterableType(expr, typeOfStarParent.iteratedItemType) + } + } + return null + } + + is PyDoubleStarExpression -> { + if (parent.parent is PyExpression) { + val typeOfDoubleStarParent = getExpectedType(parent.parent as PyExpression, ctx) + if (PyNames.MAPPING == typeOfDoubleStarParent?.name) { + return typeOfDoubleStarParent + } + if (typeOfDoubleStarParent is PyCollectionType) { + // upcast to Map + return PyCollectionTypeImpl.createTypeByQName(expr, "typing." + PyNames.MAPPING, false, typeOfDoubleStarParent.elementTypes) + } + } + return null + } + + is PyKeywordArgument -> { + if (parent.valueExpression == expr) { + return getExpectedType(parent, ctx) + } + return null + } + + is PyTupleExpression -> { + val indexOfExpr = parent.elements.indexOf(expr) + val typeOfParentTuple = getExpectedType(parent, ctx) + if (typeOfParentTuple is PyTupleType && typeOfParentTuple.elementTypes.isNotEmpty()) { + return getElementTypeAtTupleIndex(parent, typeOfParentTuple, indexOfExpr) + } + if (typeOfParentTuple is PyCollectionType) { + return typeOfParentTuple.iteratedItemType + } + return null + } + + is PySetLiteralExpression, + is PyListLiteralExpression, + -> { + val typeOfParentList = getExpectedType(parent, ctx) + if (typeOfParentList is PyCollectionType) { + return typeOfParentList.iteratedItemType + } + return null + } + + is PyKeyValueExpression -> { + if (parent.parent is PyDictLiteralExpression) { + val parentDict = parent.parent as PyDictLiteralExpression + val typeOfParentDict = getExpectedType(parentDict, ctx) + if (typeOfParentDict is PyCollectionType && typeOfParentDict.elementTypes.size == 2) { + val index = if (parent.key == expr) 0 else 1 + return typeOfParentDict.elementTypes[index] + } + if (typeOfParentDict is PyTypedDictType && parent.key is PyStringLiteralExpression) { + val argName = (parent.key as PyStringLiteralExpression).stringValue + return typeOfParentDict.getElementType(argName) + } + } + return null + } + + is PyParameterList -> { + if (expr.parent.parent is PyLambdaExpression && expr is PyParameter) { + val indexOfExpr = parent.parameters.indexOf(expr) + val typeOfParentLambda = getExpectedType(parent.parent as PyExpression, ctx) + if (typeOfParentLambda is PyCallableType) { + val parameters = typeOfParentLambda.getParameters(ctx) + if (parameters != null && indexOfExpr >= 0 && indexOfExpr < parameters.size) { + return parameters[indexOfExpr].getType(ctx) + } + } + } + return null + } + + is PyLambdaExpression -> { + val typeOfParentLambda = getExpectedType(parent, ctx) + if (typeOfParentLambda is PyCallableType) { + return typeOfParentLambda.getReturnType(ctx) + } + return null + } + } + + // Compute the expected type from a given root statement/expression + return fromArgument(expr, ctx) + ?: fromAssignment(expr) + ?: fromYield(expr, ctx) + ?: fromReturn(expr, ctx) + } + + private fun fromArgument(callArgument: PyExpression, ctx: TypeEvalContext): PyType? { + val callSite = (callArgument.parent as? PyArgumentList)?.parent as? PyCallExpression + ?: callArgument.parent as? PySubscriptionExpression + ?: return null + + val argMappings = callSite.mapArguments(PyResolveContext.defaultContext(ctx)) + val argTypes = LinkedHashSet() + + for (mapping in argMappings) { + val mappedParameters = mapping.mappedParameters + + val paramType: PyType? + if (callArgument is PyStarArgument && callSite is PyCallExpression) { + paramType = fromStarArgument(callArgument, mapping, ctx) + } + else { + val param = mappedParameters[callArgument] + ?: return null // This would be a union with Any, hence return null here already + + val paramTypeOrUnpacked = param.getArgumentType(ctx) + if (paramTypeOrUnpacked is PyUnpackedTupleType) { + // happens here: f(1, "s") for function: def f(*args: *tuple[int,str]): pass; + if (paramTypeOrUnpacked.isUnbound) { + paramType = paramTypeOrUnpacked.elementTypes.firstOrNull() + } + else { + val paramIdx = mappedParameters.keys.indexOf(callArgument) + paramType = paramTypeOrUnpacked.elementTypes.getOrElse(paramIdx) { null } + } + } + else { + paramType = paramTypeOrUnpacked + } + } + argTypes.add(substituteTypeVars(paramType, callSite, mappedParameters, ctx)) + } + + return PyUnionType.union(argTypes) + } + + private fun fromStarArgument(callArgument: PyStarArgument, mapping: PyCallExpression.PyArgumentsMapping, ctx: TypeEvalContext): PyType? { + val mappedParameters = mapping.mappedParameters + if (callArgument.isKeyword) { + val param = mappedParameters.values.firstOrNull { cp -> cp.isKeywordContainer } + if (param == null) { + // The function declares no kwargs, but the caller passed a starred expression: + // E.g.: def f(s: str, n: int) gets called f(**{"s": "foo", "n": 123}). + + val dictClass = PyBuiltinCache.getInstance(callArgument).getClass("dict") ?: return null + val fields = mutableMapOf() + for (parameter in mapping.parametersMappedToVariadicKeywordArguments) { + val name = parameter.name + if (name == null || parameter.isSelf || parameter.isPositionalContainer || parameter.isKeywordContainer) { + continue + } + fields[name] = PyTypedDictType.FieldTypeAndTotality( + value = null, // We define a schema, not a specific instance value + type = parameter.getType(ctx), + qualifiers = PyTypedDictType.TypedDictFieldQualifiers(isRequired = !parameter.hasDefaultValue()) + ) + } + + return PyTypedDictType( + name = "Parameters", + fields = fields, + dictClass = dictClass, + definitionLevel = PyTypedDictType.DefinitionLevel.INSTANCE, + ancestors = emptyList(), + declaration = mapping.callableType?.declarationElement + ) + } + else { + return param.getType(ctx) + } + } + else { + val param = mappedParameters.values.firstOrNull { cp -> cp.isPositionalContainer } + if (param == null) { + // The function declares no varargs, but the caller passed a starred expression: + // E.g.: def f(s: str, n: int) gets called f(*("foo", 123)). + + val paramTypes = mapping.parametersMappedToVariadicPositionalArguments.map { cp -> cp.getType(ctx) } + return PyTupleType.create(callArgument, paramTypes) + } + else { + return param.getType(ctx) + } + } + } + + private fun substituteTypeVars( + paramType: PyType?, + callSite: PyCallSiteExpression, + mappedParameters: Map, + ctx: TypeEvalContext, + ): PyType? { + if (!hasGenerics(paramType, ctx)) return paramType + + val receiver = callSite.getReceiver(null) + val substitutions = unifyGenericCall(receiver, mappedParameters, ctx) // might cause recursion + if (substitutions == null) return paramType + + return substitute(paramType, substitutions, ctx) + } + + private fun fromWalrus(expr: PyExpression): PyType? { + val parent = expr.parent as? PyAssignmentExpression ?: return null + if (parent.assignedValue != expr) return null + val lhs = parent.target ?: return null + val rhs = parent.assignedValue ?: return null + val avoidControlFlowCtx = TypeEvalContext.codeInsightFallback(null) + return fromLhs(lhs, rhs, avoidControlFlowCtx) + } + + private fun fromAssignment(expr: PyExpression): PyType? { + val parent = expr.parent as? PyAssignmentStatement ?: return null + if (parent.assignedValue != expr) return null + val lhs = parent.leftHandSideExpression ?: return null + val rhs = parent.assignedValue ?: return null + val avoidControlFlowCtx = TypeEvalContext.codeInsightFallback(null) + return fromLhs(lhs, rhs, avoidControlFlowCtx) + } + + private fun fromLhs(lhs: PyExpression, rhs: PyExpression?, ctx: TypeEvalContext): PyType? { + if (lhs is PyParenthesizedExpression || rhs is PyParenthesizedExpression) { + // unwrap parentheses + val lhsUnparenthesized = flattenParens(lhs) as? PyExpression ?: return null + val rhsUnparenthesized = flattenParens(rhs) as? PyExpression ?: return null + return fromLhs(lhsUnparenthesized, rhsUnparenthesized, ctx) + } + + when (lhs) { + is PySequenceExpression -> { + // try to mutually descent the nested sequences both on lhs and on rhs + val tupleElementTypes = ArrayList() + for (idx in 0 until lhs.elements.size) { + val lhsElem = lhs.elements[idx] + val rhsElem = if (rhs is PySequenceExpression) rhs.elements[idx] else null + val elemType = fromLhs(lhsElem, rhsElem, ctx) + tupleElementTypes.add(elemType) + } + if (rhs is PyTupleExpression) { + // On the RHS of the assignment we are inside a tuple, hence it is safe to downcast the current type to tuple. + // The benefit is that we can preserve the positional element type information. + return PyTupleType.create(lhs, tupleElementTypes) + } + else { + val iterableElementTypes = ArrayList() + for (tupleElementType in tupleElementTypes) { + if (tupleElementType is PyUnpackedTupleType) { + iterableElementTypes.addAll(tupleElementType.elementTypes) // simplify + } + else { + iterableElementTypes.add(tupleElementType) + } + } + if (iterableElementTypes.contains(null)) { + return createIterableType(lhs, null) // simplify + } + val iterableElementTypesUnion = PyUnionType.union(iterableElementTypes) + return createIterableType(lhs, iterableElementTypesUnion) + } + } + + is PyStarExpression -> { + val starChild = lhs.expression + val starChildType = if (starChild == null) null else fromLhs(starChild, rhs, ctx) + if (starChildType is PyTupleType) { + if (starChildType.isHomogeneous) { + return PyUnpackedTupleTypeImpl.createUnbound(starChildType.iteratedItemType) + } + else { + return PyUnpackedTupleTypeImpl.create(starChildType.elementTypes) + } + } + return starChildType + } + + is PySubscriptionExpression -> { + val operandType = ctx.getType(lhs.operand) + val iterableType = if (operandType is PyCollectionType) operandType.iteratedItemType else null + if (lhs.indexExpression is PySliceItem) { + return createIterableType(lhs, iterableType) + } + return iterableType + } + + is PyTargetExpression -> { + // the following code is supposed to only consider explicitly declared type annotations + val resolvedReference = lhs.reference.resolve() + + if (resolvedReference is PyNamedParameter) { + val parameterList = resolvedReference.parent + val indexOfExpr = (parameterList as? PyParameterList)?.parameters?.indexOf(resolvedReference) ?: -1 + val parameterListHolder = parameterList?.parent + val callableType = when (parameterListHolder) { + is PyFunction -> ctx.getType(parameterListHolder) + is PyLambdaExpression -> getExpectedType(parameterListHolder, ctx) + else -> null + } + + if (callableType is PyCallableType) { + val parameters = callableType.getParameters(ctx) + if (parameters != null && indexOfExpr >= 0 && indexOfExpr < parameters.size) { + return parameters[indexOfExpr].getType(ctx) + } + } + return null + } + if (resolvedReference is PyTypedElement) { + val pyType = PyTypingTypeProvider().getReferenceType(resolvedReference, ctx, null) + if (pyType != null) { + return pyType.get() + } + } + val pyType = PyTypingTypeProvider().getReferenceType(lhs, ctx, null) + if (pyType != null) { + return pyType.get() + } + // TODO: maybe support types from Doc-Strings using: (expr as PyTargetExpressionImpl).getTypeFromDocString() + return null + } + } + + return null + } + + private fun fromYield(expr: PyExpression, ctx: TypeEvalContext): PyType? { + val parent = expr.parent as? PyYieldExpression ?: return null + val funScope = parent.parentOfType() ?: return null + + val returnType = ctx.getReturnType(funScope) + val generatorDescriptor = PyTypingTypeProvider.GeneratorTypeDescriptor.fromGenerator(returnType) + val yieldType = generatorDescriptor?.yieldType() + if (parent.isDelegating) { + return createIterableType(expr, yieldType) + } + return yieldType + } + + private fun fromReturn(expr: PyExpression, ctx: TypeEvalContext): PyType? { + val parent = expr.parent as? PyReturnStatement ?: return null + val funScope = parent.parentOfType() ?: return null + if (funScope.annotation == null) return null // no explicit return type annotation, hence any return value is acceptable + + val returnType = ctx.getReturnType(funScope) + if (funScope.isAsync) { + return PyTypingTypeProvider.unwrapCoroutineReturnType(returnType)?.get() + } + val generatorReturnType = PyTypingTypeProvider.GeneratorTypeDescriptor.fromGenerator(returnType)?.returnType() + return generatorReturnType ?: returnType + } + + /** + * Returns the PyType of a tuple expression at the given index. + * + * By design this method computes the complete array of the tuple elements. + * The reason is that this makes it much easier to spot bugs while having only minimal impact on memory/time performance. + */ + private fun getElementTypeAtTupleIndex(tupleExpr: PyTupleExpression, tupleType: PyTupleType, indexOfExpr: Int): PyType? { + if (indexOfExpr < 0) return null + + val tupleTypeArray = arrayOfNulls(tupleExpr.elements.size) + val elementTypes = tupleType.elementTypes + val variadicRepeatCount = tupleExpr.elements.size - elementTypes.size + 1 + var arrayIdx = 0 + for (idx in 0 until tupleType.elementTypes.size) { + val elemType = tupleType.elementTypes[idx] + if (elemType is PyUnpackedTupleType && elemType.isUnbound) { + repeat(variadicRepeatCount) { + tupleTypeArray[arrayIdx++] = elemType.elementTypes.firstOrNull() + } + continue + } + if (elemType is PyTupleType && elemType.isHomogeneous) { + repeat(variadicRepeatCount) { + tupleTypeArray[arrayIdx++] = elemType.elementTypes.firstOrNull() + } + continue + } + tupleTypeArray[arrayIdx++] = elemType + } + if (indexOfExpr < tupleTypeArray.size) { + return tupleTypeArray[indexOfExpr] + } + return null + } + + private fun createIterableType(anchor: PsiElement, elementType: PyType?): PyCollectionTypeImpl? { + return PyCollectionTypeImpl.createTypeByQName(anchor, "typing." + PyNames.ITERABLE, false, listOf(elementType)) + } +} 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 9560cae6d489..97148fe996b6 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 @@ -5,7 +5,6 @@ import com.intellij.openapi.util.*; import com.intellij.psi.PsiElement; import com.intellij.psi.PsiFile; import com.intellij.psi.PsiNamedElement; -import com.intellij.psi.util.PsiTreeUtil; import com.intellij.util.ArrayUtil; import com.intellij.util.containers.ContainerUtil; import com.jetbrains.python.PyNames; @@ -1884,68 +1883,6 @@ public final class PyTypeChecker { return null; } - @ApiStatus.Internal - public static @Nullable PyType getExpectedType(@NotNull PyExpression expression, @NotNull TypeEvalContext context) { - var parent = expression.getParent(); - // Handle keyword arguments by looking at the keyword argument node instead of the expression - PsiElement callArgument = parent instanceof PyKeywordArgument kwArg ? kwArg : expression; - if (callArgument.getParent() instanceof PyArgumentList argumentList) { - var mappingResults = argumentList.getCallExpression().multiMapArguments(PyResolveContext.defaultContext(context)); - if (mappingResults.isEmpty()) return null; - var argumentMapping = mappingResults.getFirst(); - var mapped = argumentMapping.getMappedParameters().get(callArgument); - if (mapped != null) { - var expected = mapped.getType(context); - // Extract element type from *args: tuple[T, ...] - if (mapped.isPositionalContainer() && expected instanceof PyTupleType tupleType && tupleType.isHomogeneous()) { - expected = tupleType.getElementTypes().get(0); - } - // Extract value type from **kwargs: dict[str, T] - else if (mapped.isKeywordContainer() && expected instanceof PyCollectionType dictType && - PyNames.DICT.equals(dictType.getPyClass().getName())) { - expected = ContainerUtil.getOrElse(dictType.getElementTypes(), 1, null); - } - if (hasGenerics(expected, context)) { - PyExpression receiver = argumentList.getParent() instanceof PyCallExpression callExpression - ? callExpression.getReceiver(null) - : null; - final var substitutions = unifyGenericCall(receiver, argumentMapping.getMappedParameters(), context); - if (substitutions != null) { - final var substitutionsWithUnresolvedReturnGenerics = - getSubstitutionsWithUnresolvedReturnGenerics(((PyCallable)expression).getParameters(context), expected, substitutions, - context); - return substitute(expected, substitutionsWithUnresolvedReturnGenerics, context); - } - } - return expected; - } - } - // Handle unpacking in assignments: skip PyParenthesizedExpression and PyTupleExpression - else if (PsiTreeUtil.skipParentsOfType(expression, PyParenthesizedExpression.class, - PyTupleExpression.class) instanceof PyAssignmentStatement assignment) { - List> mapping = assignment.getTargetsToValuesMapping(); - Pair matchingPair = ContainerUtil.find(mapping, pair -> pair.getSecond() == expression); - if (matchingPair != null && matchingPair.getFirst() instanceof PyTargetExpression target) { - // resolve declared type - if (target.getAnnotationValue() != null) { - return context.getType(target); - } - - var result = new PyTypingTypeProvider().getReferenceType(target, context, expression); - if (result != null) { - return result.get(); - } - } - } - else if (parent instanceof PyReturnStatement) { - var scopeOwner = ScopeUtil.getScopeOwner(expression); - if (scopeOwner instanceof PyFunction function && function.getAnnotationValue() != null) { - return context.getReturnType(function); - } - } - return null; - } - @ApiStatus.Internal public static class Generics { private final @NotNull Set typeVars = new LinkedHashSet<>(); @@ -2068,7 +2005,7 @@ public final class PyTypeChecker { MatchContext(@NotNull TypeEvalContext context, @NotNull GenericSubstitutions substitutions, boolean reversedSubstitutions) { this.context = context; - this.mySubstitutions = substitutions; + mySubstitutions = substitutions; this.reversedSubstitutions = reversedSubstitutions; } diff --git a/python/testSrc/com/jetbrains/python/PyExpectedTypeJudgmentTest.kt b/python/testSrc/com/jetbrains/python/PyExpectedTypeJudgmentTest.kt new file mode 100644 index 000000000000..b8ae3979cc55 --- /dev/null +++ b/python/testSrc/com/jetbrains/python/PyExpectedTypeJudgmentTest.kt @@ -0,0 +1,1058 @@ +// 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.intellij.openapi.util.RecursionManager +import com.intellij.openapi.util.StackOverflowPreventedException +import com.jetbrains.python.documentation.PythonDocumentationProvider +import com.jetbrains.python.fixtures.PyTestCase +import com.jetbrains.python.psi.PyBinaryExpression +import com.jetbrains.python.psi.PyCallExpression +import com.jetbrains.python.psi.PyExpression +import com.jetbrains.python.psi.types.PyExpectedTypeJudgement.getExpectedType +import com.jetbrains.python.psi.types.TypeEvalContext +import junit.framework.ComparisonFailure +import junit.framework.TestCase + + +class PyExpectedTypeJudgmentTest : PyTestCase() { + + private fun doTest(expression: String, expectedType: String, text: String) { + doTest(expression, PyExpression::class.java, expectedType, text) + } + + private fun doTest(expression: String, clazz: Class, expectedType: String, text: String) { + val textIndented = text.trimIndent() + myFixture.configureByText(PythonFileType.INSTANCE, textIndented) + val expr: PyExpression = myFixture.findElementByText(expression, clazz) + + RecursionManager.assertOnRecursionPrevention(myFixture.projectDisposable) + val context = TypeEvalContext.codeAnalysis(expr.getProject(), expr.getContainingFile()) + val actual = getExpectedType(expr, context) + val actualType = PythonDocumentationProvider.getTypeName(actual, context) + TestCase.assertEquals(expectedType, actualType) + } + + + fun testParenthesisExpression() { + doTest("1", "int", """ + x : int = (1) + """) + } + + fun testWalrusExpression() { + doTest("34", "int", """ + a : int + a = (b := 34) + """) + } + + fun testWalrusInTupleExpression() { + doTest("34", "int", """ + x: int + x, y = ((y := 34), 5) + """) + } + + fun testWalrusInParentheses() { + doTest("34", "object", """ + b: object + (b := 34) + """) + } + + fun testReassignFunctionParameter() { + doTest("expr", "int", """ + def f(b: int) : + b = expr + """) + } + + fun testExpressionAssignedToSlice() { + doTest("expr", "Iterable[int]", """ + a: list[int] + a[:] = expr + """) + } + + fun testNestedExpressionAssignedToSlice() { + doTest("expr", "int", """ + a: list[int] + a[:] = (expr,) + """) + } + + fun testStartIndexInSlice() { + doTest("start", "int | None", """ + a: list[int] + a[start:] = [1] + """) + } + + fun testStopIndexInSlice() { + doTest("stop", "int | None", """ + a: list[int] + a[:stop] = [1] + """) + } + + fun testStepIndexInSlice() { + doTest("step", "int | None", """ + a: list[int] + a[::step] = [1] + """) + } + + fun testTupleAsArgument() { + doTest("(expr, \"spam\")", "Iterable[str]", """ + from typing import Iterable + + def f(xs: Iterable[str]): + ... + f((expr, "spam")) + """) + } + + fun testExpressionInsideTupleAsArgument() { + doTest("expr", "str", """ + from typing import Iterable + + def f(xs: Iterable[str]): + ... + f((expr, "spam")) + """) + } + + fun testExpressionInsideLambdaAsArgument1() { + doTest("expr", "int", """ + from typing import Callable + + def f(fn: Callable[[int], object]): + ... + f(lambda expr: {}) + """) + } + + fun testExpressionInsideLambdaAsArgument2() { + doTest("expr", "str", """ + from typing import Callable + + def f(fn: Callable[[int], str]): + ... + f(lambda x: expr) + """) + } + + fun testExpressionInsideLambdaAsUntypedArgument() { + doTest("expr", "Any", """ + def f(fn): + ... + f(lambda expr: 2) + """) + } + + fun testExpressionInsideLambdaBodyAsUntypedArgument() { + doTest("\"hello\"", "Any", """ + def f(fn): + ... + f(lambda x = 2: (x := "hello")) + """) + } + + fun testExpressionInsideLambdaBodyAsAnyTypedArgument() { + doTest("\"hello\"", "Any", """ + from typing import Callable + + def f(fn: Callable[[Any], Any]): + ... + f(lambda x = 2: (x := "hello")) + """) + } + + fun testExpressionInsideLambdaBodyAsIntTypedArgument() { + doTest("\"hello\"", "int", """ + from typing import Callable + + def f(fn: Callable[[int], Any]): + ... + f(lambda x: (x := "hello")) + """) + } + + fun testExpressionInsideLambdaBodyAsIntTypedReturn() { + doTest("\"hello\"", "int", """ + from typing import Callable + + def f(fn: Callable[[Any], int]): + ... + f(lambda x: (x := "hello")) + """) + } + + fun testExpressionAsReturnValue() { + doTest("expr", "str", """ + from typing import Iterable + + def f(xs) -> str: + return expr + """) + } + + fun testExpressionInCallTargetAsReturnValue() { + doTest("expr()", PyCallExpression::class.java, "int", """ + def main() -> int: + return expr() + """) + } + + fun testTupleAsReturnValue() { + doTest("(expr, \"spam\")", "Iterable[str]", """ + from typing import Iterable + + def f(xs) -> Iterable[str]: + return (expr, "spam") + """) + } + + fun testExpressionInsideTupleAsReturnValue() { + doTest("expr", "str", """ + from typing import Iterable + + def f(xs) -> Iterable[str]: + return (expr, "spam") + """) + } + + fun testExpressionInsideLambdaAsReturnValue1() { + doTest("expr", "int", """ + from typing import Callable + + def f() -> Callable[[int], str]: + return lambda expr: "r" + """) + } + + fun testExpressionInsideLambdaAsReturnValue2() { + doTest("expr", "str", """ + from typing import Callable + + def f() -> Callable[[int], str]: + return lambda x: expr + """) + } + + fun testExpressionInAssignmentToAttribute() { + doTest("expr", "int", """ + class A: + a : int = 1 + + A.a = expr + """) + } + + fun testExpressionInsideLambdaOfGenericFunction() { + fixme("PY-85922", StackOverflowPreventedException::class.java) { + doTest("expr", "int", """ + from collections.abc import Callable, Iterable + + def f[T](x: Iterable[T], y: Callable[[T], object]): ... + + f([1], lambda expr: ...) + """) + } + } + + fun testExpressionInsideGenericClassAsReturnValue1() { + fixme("PY-85922", StackOverflowPreventedException::class.java) { + doTest("expr", "int", """ + from typing import Callable + + class A[T]: + def f(self, fn: Callable[[T], str]) -> float: + + A[int]().f(lambda expr: "s") + """) + } + } + + fun testExpressionInsideGenericClassAsReturnValue2() { + fixme("PY-85922", StackOverflowPreventedException::class.java) { + doTest("expr", "int", """ + from typing import Callable + + class A[T]: + def f(self, fn: Callable[[str], T]) -> float: + + A[int]().f(lambda x: expr) + """) + } + } + + fun testTupleAsReturnValueNoTypeHint() { + doTest("(expr, \"spam\")", "Any", """ + def f(xs): + return (expr, "spam") + """) + } + + fun testExpressionInsideTupleAsReturnValueNoTypeHint() { + doTest("expr", "Any", """ + def f(xs): + return (expr, "spam") + """) + } + + fun testTupleAsAssignmentValue() { + doTest("(42, (expr, \"spam\"))", "tuple[Any, tuple[str, Any]]", """ + x2: str + x1, (x2, x3) = (42, (expr, "spam")) + """) + } + + fun testExpressionInsideTupleAsAssignmentValue() { + doTest("expr", "str", """ + x2: str + x1, (x2, x3) = (42, (expr, "spam")) + """) + } + + fun testTupleAsAssignmentValueNoTypeHint() { + doTest("(42, (expr, \"spam\"))", "tuple[Any, tuple[Any, Any]]", """ + x1, (x2, x3) = (42, (expr, "spam")) + """) + } + + fun testExprAsAssignmentValueNoTypeHint() { + doTest("expr", "Iterable[Any]", """ + x1, (x2, x3) = expr + """) + } + + fun testExpressionInsideTupleAsAssignmentValueNoTypeHint() { + doTest("expr", "Any", """ + x1, (x2, x3) = (42, (expr, "spam")) + """) + } + + fun testExpressionAsAssignmentValueToList() { + doTest("expr", "int", """ + x: list[int] + x[0] = expr + """) + } + + fun testExpressionInsideTupleAsAssignmentValueToList() { + doTest("expr", "str", """ + x1: bool + x2: str + x3: int + x1, [x2, x3] = (true, (expr, "spam")) + """) + } + + fun testExpressionInsideTupleAsAssignmentValueToListNoTypeHint() { + doTest("expr", "Any", """ + x1, [x2, x3] = (42, (expr, "spam")) + """) + } + + fun testExpressionAsAssignmentValueToUnwrap1() { + doTest("expr", "Iterable[int]", """ + x: int + xs: tuple[int, ...] + x, *xs = expr + """) + } + + fun testExpressionAsTupleElementToUnwrap1() { + doTest("expr", "int", """ + x: int + xs: tuple[int, ...] + x, *xs = 1, 2, expr + """) + } + + fun testExpressionAsAssignmentValueToUnwrap2() { + doTest("expr", "Iterable[int | str]", """ + x: int + xs: tuple[int, str] + x, *xs = expr + """) + } + + fun testExpressionAsTupleElementToUnwrap2() { + doTest("expr", "str", """ + x: int + xs: tuple[int, str] + x, *xs = 1, 2, expr + """) + } + + fun testExpressionAsTupleElementToUnwrap2OutOfBounds() { + doTest("expr", "Any", """ + x: int + xs: tuple[int, str] + x, *xs = 1, 2, "3", expr + """) + } + + fun testExpressionInVariadicTupleEnd1() { + doTest("expr", "str", """ + x: tuple[str, *tuple[int, ...]] = expr, 2, 3 + """) + } + + fun testExpressionInVariadicTupleEnd2() { + doTest("expr", "int", """ + x: tuple[str, *tuple[int, ...]] = "s", expr, 3 + """) + } + + fun testExpressionInVariadicTupleEnd3() { + doTest("expr", "int", """ + x: tuple[str, *tuple[int, ...]] = "s", 2, expr + """) + } + + fun testExpressionInVariadicTupleMiddle1() { + doTest("expr", "str", """ + x: tuple[str, *tuple[int, ...], float] = expr, 2, 3.14 + """) + } + + fun testExpressionInVariadicTupleMiddle2() { + doTest("expr", "int", """ + x: tuple[str, *tuple[int, ...], float] = "s", expr, 3.14 + """) + } + + fun testExpressionInVariadicTupleMiddle3() { + doTest("expr", "float", """ + x: tuple[str, *tuple[int, ...], float] = "s", 2, expr + """) + } + + fun testExpressionInVariadicTupleStart1() { + doTest("expr", "int", """ + x: tuple[*tuple[int, ...], str] = expr, 2, "s" + """) + } + + fun testExpressionInVariadicTupleStart2() { + doTest("expr", "int", """ + x: tuple[*tuple[int, ...], str] = 1, expr, "s" + """) + } + + fun testExpressionInVariadicTupleStart3() { + doTest("expr", "str", """ + x: tuple[*tuple[int, ...], str] = 1, 2, expr + """) + } + + fun testSubscriptionExpression() { + doTest("expr", "Literal[\"1\", 2, \"foo\"]", """ + from typing import Literal + + d: dict[Literal["1", 2, "foo"], str] = {} + d[expr] + """) + } + + fun testArgumentForArgs() { + doTest("expr", "str", """ + def f(*args: str): + pass + + f(expr) + """) + } + + fun testArgumentForArgsOfUnpackedTuple1() { + doTest("expr", "int", """ + def f(*args: *tuple[int]): + pass + + f(expr) + """) + } + + fun testArgumentForArgsOfUnpackedTuple2() { + doTest("expr", "Any", """ + def f(*args: *tuple[int]): + pass + + f(1, expr) + """) + } + + fun testArgumentForArgsOfUnpackedTuple3() { + doTest("expr", "str", """ + def f(*args: *tuple[int,str]): + pass + + f(1, expr) + """) + } + + fun testArgumentForArgsOfUnpackedTuple4() { + doTest("expr", "int", """ + def f(*args: *tuple[int,...]): + pass + + f(1, expr) + """) + } + + fun testArgumentValueForKwArgs() { + doTest("value", "str", """ + def f(**kwargs: str): + pass + + f(foo="value") + """) + } + + fun testArgumentKeyForKwArgs() { + doTest("foo", "str", """ + def f(**kwargs: str): + pass + + f(foo="value") + """) + } + + fun testArgumentForPlainParameter() { + doTest("expr", "str", """ + def f(x: int, y: str): + pass + + f(42, expr) + """) + } + + fun testValueForTrivialAssignment() { + doTest("expr", "str", """ + x: str = expr + """) + } + + fun testLambdaInAssignment() { + doTest("lambda p_x, p_y: p_x + p_y", "(str, int) -> int", """ + from typing import Callable + + adder: Callable[[str, int], int] = lambda p_x, p_y: p_x + p_y + """) + } + + fun testParameter1OfLambdaInAssignment() { + doTest("p_x", "str", """ + from typing import Callable + + adder: Callable[[str, int], int] = lambda p_x, p_y: p_x + p_y + """) + } + + fun testParameter2OfLambdaInAssignment() { + doTest("p_y", "int", """ + from typing import Callable + + adder: Callable[[str, int], int] = lambda p_x, p_y: p_x + p_y + """) + } + + fun testReturnOfLambdaInAssignment() { + doTest("p_x + p_y", PyBinaryExpression::class.java, "int", """ + from typing import Callable + + adder: Callable[[str, int], int] = lambda p_x, p_y: p_x + p_y + """) + } + + fun testNestedLambda() { + doTest("yy", "float", """ + from typing import Callable + + func: Callable[[int], Callable[[float], str]] = lambda xx: lambda yy: "Hi" + """) + } + + fun testLambdaInAssignment_PreviouslyTyped() { + doTest("lambda p_x, p_y: p_x + p_y", "(str, int) -> int", """ + from typing import Callable + + adder: Callable[[str, int], int] + adder = lambda p_x, p_y: p_x + p_y + """) + } + + fun testParameter1OfLambdaInAssignment_PreviouslyTyped() { + doTest("p_x", "str", """ + from typing import Callable + + adder: Callable[[str, int], int] + adder = lambda p_x, p_y: p_x + p_y + """) + } + + fun testParameter2OfLambdaInAssignment_PreviouslyTyped() { + doTest("p_y", "int", """ + from typing import Callable + + adder: Callable[[str, int], int] + adder = lambda p_x, p_y: p_x + p_y + """) + } + + fun testReturnOfLambdaInAssignment_PreviouslyTyped() { + doTest("p_x + p_y", PyBinaryExpression::class.java, "int", """ + from typing import Callable + + adder: Callable[[str, int], int] + adder = lambda p_x, p_y: p_x + p_y + """) + } + + fun testNestedLambda_PreviouslyTyped() { + doTest("yy", "float", """ + from typing import Callable + + func: Callable[[int], Callable[[float], str]] + func = lambda xx: lambda yy: "Hi" + """) + } + + fun testLambdaInAssignment_TypedAsAttribute() { + doTest("lambda p_x, p_y: p_x + p_y", "(str, int) -> int", """ + from typing import Callable + + class C: + attr: Callable[[str, int], int] + def __init__(self): + self.attr = lambda p_x, p_y: p_x + p_y + """) + } + + fun testParameter1OfLambdaInAssignment_TypedAsAttribute() { + doTest("p_x", "str", """ + from typing import Callable + + class C: + attr: Callable[[str, int], int] + def __init__(self): + self.attr = lambda p_x, p_y: p_x + p_y + """) + } + + fun testParameter2OfLambdaInAssignment_TypedAsAttribute() { + doTest("p_y", "int", """ + from typing import Callable + + class C: + attr: Callable[[str, int], int] + def __init__(self): + self.attr = lambda p_x, p_y: p_x + p_y + """) + } + + fun testReturnOfLambdaInAssignment_TypedAsAttribute() { + doTest("p_x + p_y", PyBinaryExpression::class.java, "int", """ + from typing import Callable + + class C: + attr: Callable[[str, int], int] + def __init__(self): + self.attr = lambda p_x, p_y: p_x + p_y + """) + } + + fun testNestedLambda_TypedAsAttribute() { + doTest("yy", "float", """ + from typing import Callable + + class C: + attr: Callable[[int], Callable[[float], str]] + def __init__(self): + self.attr = lambda xx: lambda yy: "Hi" + """) + } + + fun testListLiteral() { + doTest("[expr, 2, 3]", "list[int]", """ + from typing import List + + v: List[int] = [expr, 2, 3] + """) + } + + fun testExpressionInList() { + doTest("expr", "int", """ + from typing import List + + v: List[int] = [expr, 2, 3] + """) + } + + fun testStarArgumentExpressionInList() { + doTest("expr", "Iterable[int]", """ + from typing import List + + v: List[int] = [1, *expr, 4] + """) + } + + fun testExpressionInStarArgumentExpressionInList() { + doTest("expr", "int", """ + from typing import List + + v: List[int] = [1, *[expr, 3], 4] + """) + } + + fun testDictLiteralAsArgument() { + doTest("{'key': expr}", "dict[str, int]", """ + v: dict[str, int] = {'key': expr} + """) + } + + fun testDictKeyInDictLiteral() { + doTest("key", "str", """ + v: dict[str, int] = {'key': expr} + """) + } + + fun testDictValueInDictLiteral() { + doTest("value", "int", """ + v: dict[str, int] = {'key': value} + """) + } + + fun testStarArgumentInDictLiteral() { + doTest("expr", "Mapping[str, int]", """ + v: dict[str, int] = {'key': 1, **expr} + """) + } + + fun testDoubleStarExpressionOwnTypeShouldBeAny() { + doTest("**xs", "Any", """ + ys: dict[str, int] = {**xs} + """) + } + + fun testKeyInStarArgumentInDictLiteral() { + doTest("expr", "str", """ + v: dict[str, int] = {'key': 1, **{expr: 2}} + """) + } + + fun testValueInStarArgumentInDictLiteral() { + doTest("expr", "int", """ + v: dict[str, int] = {'key': 1, **{"otherKey": expr}} + """) + } + + fun testSetLiteralAsArgument() { + doTest("{expr, 2, 3}", "set[int]", """ + from typing import Set + + v: Set[int] = {expr, 2, 3} + """) + } + + fun testExpressionInsideSetAsArgument() { + doTest("expr", "int", """ + from typing import Set + + v: Set[int] = {expr, 2, 3} + """) + } + + fun testNonStarredExpressionAsArgument() { + doTest("expr", "int", """ + def f(*args: int): + pass + + f(expr) + """) + } + + fun testStarredExpressionAsArgument() { + doTest("expr", "tuple[int, ...]", """ + def f(*args: int): + pass + + f(*expr) + """) + } + + fun testStarredExpressionElementAsArgument1() { + doTest("123", "int", """ + def f(*args: int): + pass + + f(*(123, 456)) + """) + } + + fun testStarredExpressionElementAsArgument2() { + doTest("123", "int", """ + def f(s: str, *args: int): + pass + + f("foo", *(123, 456)) + """) + } + + fun testStarredExpressionElementAsArgument3() { + doTest("123", "int", """ + def f(s: str, n: int): + pass + + f(*("foo", 123)) + """) + } + + fun testDoubleStarredExpressionElementAsArgument1() { + doTest("123", "int", """ + def f(**kwargs: int): + pass + + f(**{"s": 123, "n": 456}) + """) + } + + fun testDoubleStarredExpressionElementAsArgument2() { + doTest("123", "int", """ + def f(s: str, **kwargs: int): + pass + + f("foo", **{"s2": 123, "n": 456}) + """) + } + + fun testDoubleStarredExpressionElementAsArgument3() { + doTest("123", "int", """ + def f(s: str, n: int): + pass + + f(**{"s": "foo", "n": 123}) + """) + } + + fun testDoubleStarredExpressionElementAsArgument1B() { + doTest("123", "int", """ + from typing import TypedDict, Unpack + + class FArgs(TypedDict): + s: str + n: int + + def f(**kwargs: Unpack[FArgs]): + pass + + f(**{"s": "foo", "n": 123}) + """) + } + + fun testDoubleStarredExpressionElementAsArgument2B() { + doTest("123", "int", """ + from typing import TypedDict, Unpack + + class FArgs(TypedDict): + s: str + n: int + + def f(s: str, **kwargs: Unpack[FArgs]): + pass + + f("foo", **{"s": "foo", "n": 123}) + """) + } + + fun testGenericMethodArgument() { + doTest("expr", "str", """ + class Box[T]: + def m(self, x: T): + ... + b: Box[str] + b.m(expr) + """) + } + + fun testGenericFunctionArgument() { + doTest("expr", "int", """ + def f[T](x: T, y: T) + ... + + f(42, expr) + """) + } + + fun testStarExpressionOwnTypeShouldBeInt() { + doTest("*xs", "int", """ + ys: list[int] = [1, *xs] + """) + } + + fun testStarExpressionInSetLiteral() { + doTest("xs", "Iterable[int]", """ + ys: set[int] = {1, *xs} + """) + } + + fun testStarExpressionInTupleLiteral() { + doTest("xs", "Iterable[int]", """ + ys: tuple[int, ...] = (1, *xs) + """) + } + + fun testExpressionInTupleLiteral() { + doTest("x", "int", """ + ys: tuple[int, ...] = [1, x, 3] + """) + } + + fun testExprNonDoubleStarredExpressionAsArgument() { + doTest("expr", "int", """ + def f(**kwargs: int): + pass + + f(param = expr) + """) + } + + fun testParamNonDoubleStarredExpressionAsArgument() { + doTest("param", "str", """ + def f(**kwargs: str): + pass + + f(param = expr) + """) + } + + fun testDoubleStarredExpressionAsArgument() { + doTest("**expr", "dict[str, str]", """ + def f(**kwargs: str): + pass + + f(**expr) + """) + } + + fun testDoubleStarredExpressionKeyAsArgument() { + doTest("key", "str", """ + def f(**kwargs: str): + pass + + f(**{"key" : 0}) + """) + } + + fun testDoubleStarredExpressionValueAsArgument() { + doTest("0", "str", """ + def f(**kwargs: str): + pass + + f(**{"key" : 0}) + """) + } + + fun testArgumentOfOverloadedFunctions() { + doTest("expr", "int | str", """ + from typing import overload + + @overload + def f(x: int) -> int: ... + + @overload + def f(x: str) -> str: ... + + def f(x): return x + + f(expr) + """) + } + + fun testArgumentOfOverloadedFunctionsBoundedByReturn() { + fixme("Depends on correct function overload matching", ComparisonFailure::class.java) { + doTest("expr", "str", """ + from typing import overload + + @overload + def f(x: int) -> int: ... + + @overload + def f(x: str) -> str: ... + + def f(x): return x + + a: str = f(expr) + """) + } + } + + fun testReturnOfOverloadedFunctions() { + fixme("Depends on correct function overload matching", ComparisonFailure::class.java) { + doTest("expr", "int", """ + from typing import overload + + @overload + def f(x: int) -> int: ... + + @overload + def f(x: str) -> str: ... + + def f(x): return x + + expr = f(1) + """) + } + } + + fun testReturnInAsyncFunction() { + doTest("expr", "object", """ + async def foo() -> object: + return expr + """) + } + + fun testYieldExpressionInTypedGenerator() { + doTest("send", "int", """ + from typing import Generator + + def f() -> Generator[int, str, float]: + receive = yield send + return result + """) + } + + fun testReturnInTypedGenerator() { + doTest("result", "float", """ + from typing import Generator + + def f() -> Generator[int, str, float]: + receive = yield send + return result + """) + } + + fun testYieldExpressionFromGenerator() { + doTest("expr", "Iterable[int]", """ + from typing import Generator + + def main() -> Generator[int]: + yield from expr + """) + } + + // Note: The return type of yield is not subject to the [PyExpectedTypeJudgement] computation. + @Suppress("unused") + fun do_not_testReturnFromYieldExpressionInTypedGenerator() { + doTest("receive", "Any", """ + from typing import Generator + + def f() -> Generator[int, str, float]: + receive = yield send + return result + """) + } +}