PY-87859 / PY-87913: variance error on frozen dataclass / ReadOnly attribute

(cherry picked from commit cf4316e5c0fb0c7230348d52cba7bd333d64667a)

IJ-MR-192922

GitOrigin-RevId: 0b9e9955ec63534feee9c8bbcfdb2a07f73ee95e
This commit is contained in:
Marcus Mews
2026-02-25 11:15:51 +00:00
committed by intellij-monorepo-bot
parent 27dede17a6
commit c49593a948
6 changed files with 156 additions and 117 deletions
@@ -1841,6 +1841,13 @@ class PyTypingTypeProvider : PyTypeProviderWithCustomContext<Context?>() {
}
}
@JvmStatic
fun <T> isReadOnly(owner: T, context: TypeEvalContext): Boolean where T : PyTypeCommentOwner?, T : PyAnnotationOwner? {
return PyUtil.getParameterizedCachedValue(owner!!, context) {
typeHintedWithName(owner, context, READONLY, READONLY_EXT)
}
}
@JvmStatic
fun <T> isClassVar(owner: T, context: TypeEvalContext): Boolean where T : PyAnnotationOwner?, T : PyTypeCommentOwner? {
return PyUtil.getParameterizedCachedValue(owner!!, context) {
@@ -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 {
@@ -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<PsiElement> { 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<Variance>,
currentVariance: Variance,
val context: TypeEvalContext,
) : PyTypeVisitorExt<Unit>() {
val varianceStack: MutableList<Variance> = 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)
}
}
}
}
@@ -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`() {
@@ -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
""")
}
}
@@ -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
""")
}
}