mirror of
https://gitflic.ru/project/openide/openide.git
synced 2026-09-27 10:03:11 +07:00
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
This commit is contained in:
committed by
intellij-monorepo-bot
parent
656fed0013
commit
14e295879d
+5
-6
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
+8
-28
@@ -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<CompletionPar
|
||||
override fun addCompletions(parameters: CompletionParameters, context: ProcessingContext, result: CompletionResultSet) {
|
||||
val position = parameters.position.parent as? PyExpression ?: return
|
||||
if (!(position is PyStringLiteralExpression || position is PyReferenceExpression && !position.isQualified)) return
|
||||
|
||||
val typeEvalContext = TypeEvalContext.codeCompletion(position.project, position.containingFile)
|
||||
|
||||
val mappedParameters = position.getMappedParameters(PyResolveContext.defaultContext(typeEvalContext))
|
||||
if (mappedParameters != null) {
|
||||
val types = mappedParameters.mapNotNull { it.getArgumentType(typeEvalContext) }
|
||||
addToResult(position, types, result)
|
||||
return
|
||||
}
|
||||
|
||||
val assignmentStatement = PsiTreeUtil.skipParentsOfType(position,
|
||||
PyParenthesizedExpression::class.java,
|
||||
PyTupleExpression::class.java) as? PyAssignmentStatement
|
||||
if (assignmentStatement != null) {
|
||||
val mapping = assignmentStatement.targetsToValuesMapping.find { PyPsiUtils.flattenParens(it.second) === position }
|
||||
if (mapping != null) {
|
||||
val type = typeEvalContext.getType(mapping.first)
|
||||
addToResult(position, listOfNotNull(type), result)
|
||||
}
|
||||
return
|
||||
}
|
||||
val expectedType = PyExpectedTypeJudgement.getExpectedType(position, typeEvalContext)
|
||||
addToResult(position, listOfNotNull(expectedType), result)
|
||||
}
|
||||
|
||||
private fun addToResult(position: PyExpression, possibleTypes: List<PyType>, result: CompletionResultSet) {
|
||||
|
||||
-2
@@ -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<String, PyTypedDictType.FieldTypeAndTotality>
|
||||
|
||||
+5
-2
@@ -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<PyCallableParameter>();
|
||||
|
||||
@@ -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<PyType?>()
|
||||
|
||||
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<String, PyTypedDictType.FieldTypeAndTotality>()
|
||||
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<PyExpression, PyCallableParameter>,
|
||||
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<PyType?>()
|
||||
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<PyType?>()
|
||||
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<PyFunction>() ?: 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<PyFunction>() ?: 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<PyType?>(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))
|
||||
}
|
||||
}
|
||||
@@ -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<Pair<PyExpression, PyExpression>> mapping = assignment.getTargetsToValuesMapping();
|
||||
Pair<PyExpression, PyExpression> 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<PyTypeVarType> 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;
|
||||
}
|
||||
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user