From f564bebd6061aa78580f0c516cb7ffd758837d7f Mon Sep 17 00:00:00 2001 From: Marcus Mews Date: Wed, 18 Feb 2026 10:44:42 +0000 Subject: [PATCH] PY-87731 Variance of type parameters of Protocols should be inferred as bivariant (cherry picked from commit 5b49f1420b6553966baf07c5e141b46a27061f01) IJ-MR-192193 GitOrigin-RevId: 3dcb42e2736ceb269ce50df0088742775d8c76e5 --- .../inspections/PyVarianceInspection.kt | 1 + .../psi/types/PyExpectedVarianceJudgment.kt | 14 +++++----- .../python/PyExpectedVarianceJudgmentTest.kt | 26 ++++++++++++++++--- .../python/PyInferredVarianceJudgmentTest.kt | 20 ++++++++++++++ 4 files changed, 50 insertions(+), 11 deletions(-) diff --git a/python/python-psi-impl/src/com/jetbrains/python/inspections/PyVarianceInspection.kt b/python/python-psi-impl/src/com/jetbrains/python/inspections/PyVarianceInspection.kt index 9c535983586d..cd0d0520a9a4 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/inspections/PyVarianceInspection.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/inspections/PyVarianceInspection.kt @@ -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)) { 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 290d8fda5a1c..2ccbb724930a 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 @@ -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 diff --git a/python/testSrc/com/jetbrains/python/PyExpectedVarianceJudgmentTest.kt b/python/testSrc/com/jetbrains/python/PyExpectedVarianceJudgmentTest.kt index 9d12a7703fcb..0ca2cd04c5ec 100644 --- a/python/testSrc/com/jetbrains/python/PyExpectedVarianceJudgmentTest.kt +++ b/python/testSrc/com/jetbrains/python/PyExpectedVarianceJudgmentTest.kt @@ -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]: diff --git a/python/testSrc/com/jetbrains/python/PyInferredVarianceJudgmentTest.kt b/python/testSrc/com/jetbrains/python/PyInferredVarianceJudgmentTest.kt index 1f3baa3d3ed2..1b1a34d4e5ea 100644 --- a/python/testSrc/com/jetbrains/python/PyInferredVarianceJudgmentTest.kt +++ b/python/testSrc/com/jetbrains/python/PyInferredVarianceJudgmentTest.kt @@ -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]: