PY-87731 Variance of type parameters of Protocols should be inferred as bivariant

(cherry picked from commit 5b49f1420b6553966baf07c5e141b46a27061f01)

IJ-MR-192193

GitOrigin-RevId: 3dcb42e2736ceb269ce50df0088742775d8c76e5
This commit is contained in:
Marcus Mews
2026-02-19 08:36:01 +00:00
committed by intellij-monorepo-bot
parent 538a8c9c3e
commit f564bebd60
4 changed files with 50 additions and 11 deletions
@@ -127,6 +127,7 @@ class PyVarianceInspection : PyInspection() {
context: TypeEvalContext,
) {
val varianceExpected = PyExpectedVarianceJudgment.getExpectedVariance(node, context) ?: return
if (varianceExpected == BIVARIANT) return
val varianceInferred = PyInferredVarianceJudgment.getInferredVariance(node, context) ?: return
if (!isCompatibleWith(varianceInferred, varianceExpected)) {
@@ -2,6 +2,9 @@ package com.jetbrains.python.psi.types
import com.intellij.psi.PsiElement
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.psi.PyAnnotation
import com.jetbrains.python.psi.PyAnnotationOwner
import com.jetbrains.python.psi.PyArgumentList
@@ -63,9 +66,9 @@ object PyExpectedVarianceJudgment {
-> {
when (parent) {
is PySubscriptionExpression,
-> fromElementInSubscriptionExpression(element, 0, parent, context)
-> fromElementInSubscriptionExpression(0, parent, context)
is PyTupleExpression if parent.parent is PySubscriptionExpression
-> fromElementInSubscriptionExpression(element, parent.elements.indexOf(element),
-> fromElementInSubscriptionExpression(parent.elements.indexOf(element),
parent.parent as PySubscriptionExpression, context)
else
-> getExpectedVariance(parent, context)
@@ -89,7 +92,6 @@ object PyExpectedVarianceJudgment {
}
private fun fromElementInSubscriptionExpression(
element: PsiElement,
refIndex: Int,
subscriptionExpr: PySubscriptionExpression,
context: TypeEvalContext,
@@ -101,10 +103,8 @@ object PyExpectedVarianceJudgment {
qualifierType = PyTypeChecker.findGenericDefinitionType(qualifierType.pyClass, context) ?: qualifierType
}
if (qualifierType is PyClassLikeType && qualifierType.classQName == PyTypingTypeProvider.GENERIC) {
val refExpr = element as? PyReferenceExpression ?: return null
val typeVarType = PyTypingTypeProvider.getType(refExpr, context)?.get() as? PyTypeVarType ?: return null
return getInferredVariance(typeVarType, context)
if (qualifierType is PyClassLikeType && setOf(GENERIC, PROTOCOL, PROTOCOL_EXT).contains(qualifierType.classQName)) {
return BIVARIANT // for T in: `class C(Generic[T])` or `class C(Protocol[T])`
}
if (qualifierType is PyCollectionType) {
val paramVariance = getTypeParameterVarianceAtIndex(qualifierType, refIndex, context) ?: return null
@@ -20,8 +20,17 @@ internal class PyExpectedVarianceJudgmentTest : PyTestCase() {
assertEquals(expectedVariance, actualVariance)
}
fun `test Generic class legacy syntax 1co`() {
doTest("T1]", Variance.COVARIANT, """
fun `test Generic super class expects bivariant type parameters`() {
doTest("T]", Variance.BIVARIANT, """
from typing import TypeVar, Generic
T = TypeVar("T")
class C(Generic[T]):
pass
""")
}
fun `test Generic super class expects bivariant type parameters co`() {
doTest("T1]", Variance.BIVARIANT, """
from typing import TypeVar, Generic
T1 = TypeVar("T1", covariant=True)
class Box(Generic[T1]):
@@ -29,8 +38,8 @@ internal class PyExpectedVarianceJudgmentTest : PyTestCase() {
""")
}
fun `test Generic class legacy syntax 1contra`() {
doTest("T1]", Variance.CONTRAVARIANT, """
fun `test Generic super class expects bivariant type parameters contra`() {
doTest("T1]", Variance.BIVARIANT, """
from typing import TypeVar, Generic
T1 = TypeVar("T1", contravariant=True)
class Box(Generic[T1]):
@@ -38,6 +47,15 @@ internal class PyExpectedVarianceJudgmentTest : PyTestCase() {
""")
}
fun `test Protocol super class expects bivariant type parameters`() {
doTest("T]", Variance.BIVARIANT, """
from typing import TypeVar, Protocol
T = TypeVar("T")
class C(Protocol[T]):
pass
""")
}
fun `test Generic class attribute`() {
doTest("T #", Variance.INVARIANT, """
class A[T]:
@@ -139,6 +139,26 @@ internal class PyInferredVarianceJudgmentTest : PyTestCase() {
""")
}
fun `test Generic class unused TypeVar syntax`() {
// see comment about bivariance in: PyInferredVarianceJudgment.doGetInferredVariance
doTest("T])", Variance.BIVARIANT, PyReferenceExpression::class.java, """
from typing import TypeVar, Generic
T = TypeVar("T", infer_variance=True)
class A(Generic[T]):
def method(self): pass
""")
}
fun `test Generic protocol class unused TypeVar syntax`() {
// see comment about bivariance in: PyInferredVarianceJudgment.doGetInferredVariance
doTest("T])", Variance.BIVARIANT, PyReferenceExpression::class.java, """
from typing import TypeVar, Protocol
T = TypeVar("T", infer_variance=True)
class A(Protocol[T]):
def method(self): pass
""")
}
fun `test Generic class attribute`() {
doTest("T", Variance.INVARIANT, """
class A[T]: