diff --git a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.kt b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.kt index 2b4e1003ae44..291b113e1a7d 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.kt @@ -1841,6 +1841,13 @@ class PyTypingTypeProvider : PyTypeProviderWithCustomContext() { } } + @JvmStatic + fun isReadOnly(owner: T, context: TypeEvalContext): Boolean where T : PyTypeCommentOwner?, T : PyAnnotationOwner? { + return PyUtil.getParameterizedCachedValue(owner!!, context) { + typeHintedWithName(owner, context, READONLY, READONLY_EXT) + } + } + @JvmStatic fun isClassVar(owner: T, context: TypeEvalContext): Boolean where T : PyAnnotationOwner?, T : PyTypeCommentOwner? { return PyUtil.getParameterizedCachedValue(owner!!, context) { diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyExpectedVarianceJudgment.kt b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyExpectedVarianceJudgment.kt index 2ccbb724930a..7682295d7f7f 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyExpectedVarianceJudgment.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyExpectedVarianceJudgment.kt @@ -1,23 +1,30 @@ package com.jetbrains.python.psi.types import com.intellij.psi.PsiElement +import com.jetbrains.python.codeInsight.parseStdDataclassParameters import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider.Companion.GENERIC import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider.Companion.PROTOCOL import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider.Companion.PROTOCOL_EXT +import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider.Companion.isFinal +import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider.Companion.isReadOnly import com.jetbrains.python.psi.PyAnnotation import com.jetbrains.python.psi.PyAnnotationOwner import com.jetbrains.python.psi.PyArgumentList +import com.jetbrains.python.psi.PyBinaryExpression import com.jetbrains.python.psi.PyClass +import com.jetbrains.python.psi.PyExpressionStatement import com.jetbrains.python.psi.PyFunction import com.jetbrains.python.psi.PyListLiteralExpression import com.jetbrains.python.psi.PyNamedParameter import com.jetbrains.python.psi.PyParameterList import com.jetbrains.python.psi.PyReferenceExpression import com.jetbrains.python.psi.PyStatementList +import com.jetbrains.python.psi.PyStringLiteralExpression import com.jetbrains.python.psi.PySubscriptionExpression import com.jetbrains.python.psi.PyTargetExpression import com.jetbrains.python.psi.PyTupleExpression +import com.jetbrains.python.psi.PyTypeAliasStatement import com.jetbrains.python.psi.PyTypeCommentOwner import com.jetbrains.python.psi.PyTypeDeclarationStatement import com.jetbrains.python.psi.types.PyInferredVarianceJudgment.attributeDoesNotAffectVarianceInference @@ -45,9 +52,11 @@ object PyExpectedVarianceJudgment { return when (element) { is PyClass, - -> COVARIANT + is PyTypeAliasStatement, + is PyExpressionStatement, // parent of synthetic expressions created by PyElementGenerator#createExpressionFromText() + -> BIVARIANT is PyFunction, - -> fromFunction(element, parent, context) + -> fromFunction(element, parent) is PyTypeDeclarationStatement, -> fromTypeDeclarationStatement(element, parent, context) is PyNamedParameter, @@ -57,9 +66,11 @@ object PyExpectedVarianceJudgment { // keep the following list as precise and short as possible to enforce returning null whenever possible is PyArgumentList, + is PyBinaryExpression, is PyParameterList, is PyStatementList, is PyAnnotation, + is PyStringLiteralExpression, is PyReferenceExpression, is PySubscriptionExpression, is PyTupleExpression, @@ -79,16 +90,18 @@ object PyExpectedVarianceJudgment { } } - private fun fromFunction(function: PyFunction, parent: PsiElement, context: TypeEvalContext): Variance? { + private fun fromFunction(function: PyFunction, parent: PsiElement): Variance? { + if (parent !is PyStatementList && parent.parent !is PyClass) return null if (functionDoesNotAffectVarianceInference(function)) return null - return getExpectedVariance(parent, context) + return COVARIANT } private fun fromTypeDeclarationStatement(element: PyTypeDeclarationStatement, parent: PsiElement, context: TypeEvalContext): Variance? { - if (parent.parent !is PyClass) return null + val parentClass = parent.parent as? PyClass ?: return null val targetExpr = element.target as? PyTargetExpression ?: return null if (attributeDoesNotAffectVarianceInference(targetExpr)) return null - return if (isFinal(targetExpr, context)) COVARIANT else INVARIANT + if (isEffectivelyReadOnly(targetExpr, parentClass, context)) return COVARIANT + return INVARIANT } private fun fromElementInSubscriptionExpression( @@ -141,11 +154,14 @@ object PyExpectedVarianceJudgment { return type is PyClassLikeType && PyTypingTypeProvider.CALLABLE == type.classQName } - private fun isFinal(element: PsiElement, context: TypeEvalContext): Boolean { + private fun isEffectivelyReadOnly(element: PsiElement, parentClass: PyClass, context: TypeEvalContext): Boolean { if (element is PyTypeCommentOwner && element is PyAnnotationOwner) { - return PyTypingTypeProvider.isFinal(element, context) + if (isFinal(element, context) || isReadOnly(element, context)) { + return true + } } - return false + val isFrozen = parseStdDataclassParameters(parentClass, context)?.frozen ?: false + return isFrozen } private fun Variance.invert(): Variance { diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyInferredVarianceJudgment.kt b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyInferredVarianceJudgment.kt index 35fe8634c93e..b43ca78276df 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyInferredVarianceJudgment.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyInferredVarianceJudgment.kt @@ -6,13 +6,16 @@ import com.intellij.psi.PsiElement import com.intellij.psi.util.PsiTreeUtil import com.intellij.util.Processor import com.jetbrains.python.PyNames -import com.jetbrains.python.codeInsight.parseStdDataclassParameters import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider +import com.jetbrains.python.psi.LanguageLevel +import com.jetbrains.python.psi.PyAnnotation import com.jetbrains.python.psi.PyCallExpression import com.jetbrains.python.psi.PyClass +import com.jetbrains.python.psi.PyElementGenerator import com.jetbrains.python.psi.PyFunction import com.jetbrains.python.psi.PyQualifiedNameOwner import com.jetbrains.python.psi.PyReferenceExpression +import com.jetbrains.python.psi.PyStringLiteralExpression import com.jetbrains.python.psi.PyTargetExpression import com.jetbrains.python.psi.PyTypeAliasStatement import com.jetbrains.python.psi.PyTypeParameter @@ -92,8 +95,7 @@ object PyInferredVarianceJudgment { is PyClass -> collector.collectInClass(tvId, tvId.scopeOwner) is PyTypeAliasStatement -> { val typeExpression = tvId.scopeOwner.typeExpression ?: return INVARIANT - val typeAliasType = PyTypingTypeProvider.getType(typeExpression, context)?.get() ?: return INVARIANT - collector.collectInType(tvId, typeAliasType, COVARIANT) + collector.collectReferencesTo(tvId, typeExpression) } else -> return INVARIANT } @@ -120,10 +122,9 @@ object PyInferredVarianceJudgment { fun collectInClass(tvId: TypeVariableId, clazz: PyClass) { val classType = context.getType(clazz) if (classType is PyClassLikeType) { - val isDataclassFrozen = parseStdDataclassParameters(clazz, context)?.frozen ?: false val processor = Processor { element -> when (element) { - is PyTargetExpression -> collectInAttribute(tvId, element, isDataclassFrozen) + is PyTargetExpression -> collectInAttribute(tvId, element) is PyFunction -> collectInFunction(tvId, element) } val isInvariantAlready = usages.contains(INVARIANT) || (usages.contains(COVARIANT) && usages.contains(CONTRAVARIANT)) @@ -138,116 +139,44 @@ object PyInferredVarianceJudgment { for (superClassExpr in clazz.superClassExpressions) { val superType = PyTypingTypeProvider.getType(superClassExpr, context)?.get() ?: continue if (superType !is PyCollectionType) continue - val declaredType = PyTypeChecker.findGenericDefinitionType(superType.pyClass, context) ?: continue - - val typeParams = declaredType.elementTypes - val typeArgs = superType.elementTypes - val idxMax = typeParams.size.coerceAtMost(typeArgs.size) - for (idx in 0 until idxMax) { - val typeParam = typeParams[idx] as? PyTypeVarType ?: continue - val typeArg = superType.elementTypes[idx] - if (typeArg is PyTypeVarType && typeArg.name == tvId.name && typeArg.scopeOwner == tvId.scopeOwner && typeParam.scopeOwner != null) { - // we need to infer the variance of the base classes type variable: `class Derived[T](Base[T])` - val substitutedTvId = TypeVariableId(typeParam.name, typeParam.scopeOwner!!) - collectInType(substitutedTvId, declaredType, BIVARIANT) - } - } + collectReferencesTo(tvId, superClassExpr) } } - private fun collectInAttribute(tvId: TypeVariableId, target: PyTargetExpression, isDataclassFrozen: Boolean) { + private fun collectInAttribute(tvId: TypeVariableId, target: PyTargetExpression) { if (attributeDoesNotAffectVarianceInference(target)) return - val attributeType = context.getType(target) - val variance = if (isDataclassFrozen || PyTypingTypeProvider.isFinal(target, context)) COVARIANT else INVARIANT - collectInType(tvId, attributeType, variance) + collectReferencesTo(tvId, target.annotation) } private fun collectInFunction(tvId: TypeVariableId, function: PyFunction) { if (functionDoesNotAffectVarianceInference(function)) return val callableType = context.getType(function) as? PyCallableType ?: return - - val returnType = callableType.getReturnType(context) - collectInType(tvId, returnType, COVARIANT) + collectReferencesTo(tvId, function.annotation) val parameters = callableType.getParameters(context) ?: return for (parameter in parameters) { if (parameter.isSelf) continue - val parameterType = parameter.getType(context) - collectInType(tvId, parameterType, CONTRAVARIANT) + collectReferencesTo(tvId, parameter.parameter) } } - fun collectInType(tvId: TypeVariableId, type: PyType?, currentVariance: Variance) { - if (type == null) return - val visitor = VarianceInferenceTypeVisitor(tvId, usages, currentVariance, context) - PyTypeVisitor.visit(type, visitor) - } - } - - - private class VarianceInferenceTypeVisitor( - val tvId: TypeVariableId, - val usages: MutableSet, - currentVariance: Variance, - val context: TypeEvalContext, - ) : PyTypeVisitorExt() { - - val varianceStack: MutableList = mutableListOf(currentVariance) - - override fun visitPyTypeVarType(typeVarType: PyTypeVarType) { - if (typeVarType.name == tvId.name && typeVarType.scopeOwner == tvId.scopeOwner) { - usages.add(varianceStack.last()) + fun collectReferencesTo(tvId: TypeVariableId, element: PsiElement?) { + if (element == null) return + val annValue = if (element is PyAnnotation) element.value else element + if (annValue is PyStringLiteralExpression) { + val elementGenerator = PyElementGenerator.getInstance(annValue.project) + val syntheticElement = elementGenerator.createExpressionFromText(LanguageLevel.forElement(annValue), annValue.stringValue) + return collectReferencesTo(tvId, syntheticElement) } - } - override fun visitPyGenericType(genericType: PyCollectionType) { - val declaredType = PyTypeChecker.findGenericDefinitionType(genericType.pyClass, context) ?: return - val typeParams = declaredType.elementTypes - val typeArgs = genericType.elementTypes - val idxMax = typeParams.size.coerceAtMost(typeArgs.size) - for (idx in 0 until idxMax) { - val declaredTV = typeParams[idx] as? PyTypeVarType ?: continue - val paramVariance = if (declaredTV.variance == INFER_VARIANCE) getInferredVariance(declaredTV, context) else declaredTV.variance - val nextVariance = combineVariance(varianceStack.last(), paramVariance) - - varianceStack.add(nextVariance) - visit(typeArgs[idx], this) - varianceStack.removeAt(varianceStack.size - 1) - } - } - - override fun visitPyCallableType(callableType: PyCallableType) { - val returnType = callableType.getReturnType(context) - val nextCovariant = combineVariance(varianceStack.last(), COVARIANT) - varianceStack.add(nextCovariant) - visit(returnType, this) - varianceStack.removeAt(varianceStack.size - 1) - - val parameters = callableType.getParameters(context) ?: return - val nextContravariant = combineVariance(varianceStack.last(), CONTRAVARIANT) - for (parameter in parameters) { - val parameterType = parameter.getType(context) - varianceStack.add(nextContravariant) - visit(parameterType, this) - varianceStack.removeAt(varianceStack.size - 1) - } - } - - override fun visitPyUnionType(unionType: PyUnionType) { - for (member in unionType.members) { - visit(member, this) - } - } - - override fun visitPyIntersectionType(unionType: PyIntersectionType) { - for (member in unionType.members) { - visit(member, this) - } - } - - override fun visitPyTupleType(tupleType: PyTupleType) { - for (elementType in tupleType.elementTypes) { - visit(elementType, this) + val refExpressions = PsiTreeUtil.findChildrenOfType(element, PyReferenceExpression::class.java) + for (refExpression in refExpressions) { + val refType = PyTypingTypeProvider.getType(refExpression, context) ?: continue + val typeVarType = refType.get() as? PyTypeVarType ?: continue + if (typeVarType.name == tvId.name && typeVarType.scopeOwner == tvId.scopeOwner) { + val exprVariance = PyExpectedVarianceJudgment.getExpectedVariance(refExpression, context) ?: continue + usages.add(exprVariance) + } } } } diff --git a/python/testSrc/com/jetbrains/python/PyExpectedVarianceJudgmentTest.kt b/python/testSrc/com/jetbrains/python/PyExpectedVarianceJudgmentTest.kt index 0ca2cd04c5ec..7f023990a21e 100644 --- a/python/testSrc/com/jetbrains/python/PyExpectedVarianceJudgmentTest.kt +++ b/python/testSrc/com/jetbrains/python/PyExpectedVarianceJudgmentTest.kt @@ -6,6 +6,7 @@ import com.jetbrains.python.psi.PyExpression import com.jetbrains.python.psi.types.PyExpectedVarianceJudgment.getExpectedVariance import com.jetbrains.python.psi.types.PyTypeVarType.Variance import com.jetbrains.python.psi.types.TypeEvalContext +import junit.framework.AssertionFailedError import org.intellij.lang.annotations.Language internal class PyExpectedVarianceJudgmentTest : PyTestCase() { @@ -79,6 +80,14 @@ internal class PyExpectedVarianceJudgmentTest : PyTestCase() { """) } + fun `test Generic class readonly attribute`() { + doTest("T] #", Variance.COVARIANT, """ + from typing import ReadOnly + class A[T]: + attr: ReadOnly[T] # attribute + """) + } + fun `test Generic class final attribute`() { doTest("T] #", Variance.COVARIANT, """ from typing import Final @@ -224,9 +233,9 @@ internal class PyExpectedVarianceJudgmentTest : PyTestCase() { } fun `test Generic class type argument PEP695 syntax`() { - doTest("T2]", Variance.COVARIANT, """ + doTest("T2]", Variance.BIVARIANT, """ from typing import TypeVar, Generic - class Box[T1]: # actually bivariant, but we use covariant as a compromise + class Box[T1]: pass T2 = TypeVar('T2', contravariant=True) class ReadOnlyBox(Box[T2], Generic[T2]): @@ -235,9 +244,9 @@ internal class PyExpectedVarianceJudgmentTest : PyTestCase() { } fun `test Generic class type argument PEP695 syntax 2a`() { - doTest("T3,", Variance.COVARIANT, """ + doTest("T3,", Variance.BIVARIANT, """ from typing import TypeVar, Generic - class Box[T1, T2]: # actually bivariant, but we use covariant as a compromise + class Box[T1, T2]: pass T3 = TypeVar("T3", contravariant=True) @@ -248,9 +257,9 @@ internal class PyExpectedVarianceJudgmentTest : PyTestCase() { } fun `test Generic class type argument PEP695 syntax 2b`() { - doTest("T4]", Variance.COVARIANT, """ + doTest("T4]", Variance.BIVARIANT, """ from typing import TypeVar, Generic - class Box[T1, T2]: # actually bivariant, but we use covariant as a compromise + class Box[T1, T2]: pass T3 = TypeVar("T3", contravariant=True) @@ -305,6 +314,33 @@ internal class PyExpectedVarianceJudgmentTest : PyTestCase() { """) } + fun `test Frozen attribute`() { + doTest("T #", Variance.COVARIANT, """ + from dataclasses import dataclass + @dataclass(frozen=True) + class A[T]: + attr: T # read-only + """) + } + + fun `test String literal type`() { + doTest("T\" #", Variance.COVARIANT, """ + from dataclasses import dataclass + @dataclass(frozen=True) + class A[T]: + attr: "T" # read-only + """) + } + + fun `test String literal type at return`() { + fixme("PY-87942: No AST in string literal of type annotation", AssertionFailedError::class.java) { + doTest("T\"", Variance.COVARIANT, """ + class A[T]: + def f(self, t: Callable[["T"],None]) : ... + """) + } + } + // Expect null to avoid variance compatibility inspection check fun `test Type alias for generic class`() { diff --git a/python/testSrc/com/jetbrains/python/PyInferredVarianceJudgmentTest.kt b/python/testSrc/com/jetbrains/python/PyInferredVarianceJudgmentTest.kt index 00a0e59edf3d..6f557e53ef17 100644 --- a/python/testSrc/com/jetbrains/python/PyInferredVarianceJudgmentTest.kt +++ b/python/testSrc/com/jetbrains/python/PyInferredVarianceJudgmentTest.kt @@ -8,6 +8,7 @@ import com.jetbrains.python.psi.PyTypeParameter import com.jetbrains.python.psi.types.PyInferredVarianceJudgment.getInferredVariance import com.jetbrains.python.psi.types.PyTypeVarType.Variance import com.jetbrains.python.psi.types.TypeEvalContext +import junit.framework.AssertionFailedError import org.intellij.lang.annotations.Language internal class PyInferredVarianceJudgmentTest : PyTestCase() { @@ -192,6 +193,14 @@ internal class PyInferredVarianceJudgmentTest : PyTestCase() { """) } + fun `test Generic class readonly attribute`() { + doTest("T", Variance.COVARIANT, """ + from typing import ReadOnly + class A[T]: + attr: ReadOnly[T] # attribute + """) + } + fun `test Generic class final attribute`() { doTest("T", Variance.COVARIANT, """ from typing import Final @@ -706,11 +715,11 @@ internal class PyInferredVarianceJudgmentTest : PyTestCase() { fun `test Alias to union of class`() { doTest("U]", Variance.INVARIANT, """ - class A[T]: - def f(self, t: T): pass + class A[S]: + def f(self, t: S): pass class B[T]: def f(self) -> T: pass - type B[U] = A[U] | B[U] + type C[U] = A[U] | B[U] """) } @@ -734,13 +743,32 @@ internal class PyInferredVarianceJudgmentTest : PyTestCase() { """) } + fun `test Type in string literal`() { + fixme("PY-87942: No AST in string literal of type annotation", AssertionFailedError::class.java) { + doTest("T", Variance.COVARIANT, """ + class A[T]: + def method(self) -> "T": pass + """) + } + } + + fun `test Type in string literal with Callable`() { + fixme("PY-87942: No AST in string literal of type annotation", AssertionFailedError::class.java) { + doTest("T", Variance.COVARIANT, """ + from typing import Callable + class A[T]: + def method(self, arg: "Callable[[T], None]"): pass + """) + } + } + fun `test Recursive generic classes`() { doTest("T", Variance.COVARIANT, """ class A[T]: - def method(self) -> "B[T]": pass + def method(self) -> B[T]: pass class B[U]: - def method(self) -> "A[U]": pass + def method(self) -> A[U]: pass """) } } diff --git a/python/testSrc/com/jetbrains/python/PyVarianceTest.kt b/python/testSrc/com/jetbrains/python/PyVarianceTest.kt index 7f06f79a8a87..4f43252a9968 100644 --- a/python/testSrc/com/jetbrains/python/PyVarianceTest.kt +++ b/python/testSrc/com/jetbrains/python/PyVarianceTest.kt @@ -167,4 +167,27 @@ internal class PyVarianceTest : PyTestCase() { ... """) } + + @TestFor(issues = ["PY-87859"]) + fun `test Variance error on frozen dataclass`() { + doTestByText(""" + from dataclasses import dataclass + + @dataclass(frozen=True) + class A[T]: + a: T # expect no error here + """) + } + + @TestFor(issues = ["PY-87913"]) + fun `test Variance ReadOnly attribute`() { + doTestByText(""" + from typing import ReadOnly, TypeVar, TypedDict, Generic + + out_T = TypeVar("out_T", covariant=True) + + class TD(TypedDict, Generic[out_T]): + t: ReadOnly[out_T] # expect no error here + """) + } }