PY-57621 tuple types: widen for generic calls

(cherry picked from commit 4b8536f547d3e86dec3e07755751640d67f01f63)

GitOrigin-RevId: bb449c274a0a207e0a11eaf2112b26e458478447
This commit is contained in:
Morgan Bartholomew
2026-03-25 07:22:50 +00:00
committed by intellij-monorepo-bot
parent e6495387ce
commit b3b8047d22
9 changed files with 86 additions and 19 deletions
@@ -59,6 +59,7 @@ import java.util.stream.Stream;
import static com.jetbrains.python.PyNames.TYPE_ENUM_FLAG;
import static com.jetbrains.python.psi.PyUtil.as;
import static com.jetbrains.python.psi.types.PyTypeUtilKt.widenTupleLiterals;
public final class PyStdlibTypeProvider extends PyTypeProviderBase {
@@ -243,14 +244,8 @@ public final class PyStdlibTypeProvider extends PyTypeProviderBase {
PyExpression value = targetExpression.findAssignedValue();
if (value == null) return null;
PyType type = context.getType(value);
if (type instanceof PyTupleType tupleType) {
// until heterogeneous enums are supported, we must widen tuple types
type = PyTupleType.create(
tupleType.getDeclarationElement(),
ContainerUtil.map(tupleType.getElementTypes(), t -> PyLiteralType.upcastLiteralToClass(t))
);
}
// until heterogeneous enums are supported, we must widen tuple types
PyType type = widenTupleLiterals(context.getType(value));
return getEnumAttributeInfo(enumClass, type, context);
}
else {
@@ -60,6 +60,7 @@ import com.jetbrains.python.psi.types.PyStructuralType;
import com.jetbrains.python.psi.types.PyTupleType;
import com.jetbrains.python.psi.types.PyType;
import com.jetbrains.python.psi.types.PyTypeUtil;
import com.jetbrains.python.psi.types.PyTypeUtilKt;
import com.jetbrains.python.psi.types.PyUnionType;
import com.jetbrains.python.psi.types.TypeEvalContext;
import one.util.streamex.StreamEx;
@@ -251,7 +252,7 @@ public class PyNamedParameterImpl extends PyBaseElementImpl<PyNamedParameterStub
if (context.maySwitchToAST(this)) {
final PyExpression defaultValue = getDefaultValue();
if (defaultValue != null) {
final PyType type = PyLiteralType.upcastLiteralToClass(context.getType(defaultValue));
final PyType type = PyLiteralType.upcastLiteralToClass(PyTypeUtilKt.widenTupleLiterals(context.getType(defaultValue)));
if (type != null && !isNoneType(type)) {
if (type instanceof PyTupleType) {
return PyUnionType.createWeakType(type);
@@ -31,7 +31,8 @@ object PyCollectionTypeUtil {
private fun getListOrSetIteratedValueType(sequence: PySequenceExpression, context: TypeEvalContext): PyType? {
val elements = sequence.elements
val analyzedElementsType = PyUnionType.union(
elements.take(MAX_ANALYZED_ELEMENTS_OF_LITERALS).map { PyLiteralType.upcastLiteralToClass(context.getType(it)) }
elements.take(MAX_ANALYZED_ELEMENTS_OF_LITERALS)
.map { PyLiteralType.upcastLiteralToClass(context.getType(it).widenTupleLiterals()) }
)
return if (elements.size > MAX_ANALYZED_ELEMENTS_OF_LITERALS) {
PyUnionType.createWeakType(analyzedElementsType)
@@ -87,8 +88,8 @@ object PyCollectionTypeUtil {
.forEach {
val type = context.getType(it)
val (keyType, valueType) = getKeyValueType(type)
keyTypes.add(PyLiteralType.upcastLiteralToClass(keyType))
valueTypes.add(PyLiteralType.upcastLiteralToClass(valueType))
keyTypes.add(PyLiteralType.upcastLiteralToClass(keyType.widenTupleLiterals()))
valueTypes.add(PyLiteralType.upcastLiteralToClass(valueType.widenTupleLiterals()))
}
if (elements.size > MAX_ANALYZED_ELEMENTS_OF_LITERALS) {
@@ -1473,7 +1473,7 @@ object PyTypeChecker {
return PyCollectionTypeImpl(
genericType.pyClass, genericType.isDefinition,
genericType.elementTypes.flatMap {
flattenUnpackedTuple(clone(it))
flattenUnpackedTuple(clone<PyType>(it).widenTupleLiterals())
}
)
}
@@ -373,3 +373,15 @@ val PyType?.isUnknown: Boolean
@ApiStatus.Internal
fun PyExpression.getLiteralType(context: TypeEvalContext): PyType? =
PyLiteralType.getLiteralType(this, context)
/**
* Widens literal types within a tuple type.
* When a tuple appears nested in a non-tuple container type (e.g., `list[tuple[Literal[1], Literal["a"]]]`),
* its literal element types should be widened to their base types (e.g., `list[tuple[int, str]]`).
*/
@ApiStatus.Experimental
fun PyType?.widenTupleLiterals(): PyType? {
if (this !is PyTupleType) return this
val widenedElements = this.elementTypes.map { PyLiteralType.upcastLiteralToClass(it.widenTupleLiterals()) }
return PyTupleType(this.pyClass, widenedElements, this.isHomogeneous)
}
@@ -5262,6 +5262,48 @@ public class Py3TypeTest extends PyTestCase {
""");
}
@TestFor(issues="PY-57621")
public void testTupleInListWidens() {
doTest("list[tuple[int, str]]", """
t = (1, 'hello')
expr = [t]
""");
}
@TestFor(issues="PY-57621")
public void testTupleInTupleIsLiteral() {
var t = "tuple[Literal[1], Literal['hello']]";
doTest("tuple[" + t + ", " + t + "]", """
t = (1, 'hello')
expr = (t, t)
""");
}
@TestFor(issues="PY-57621")
public void testTupleInGenericWidens() {
doTest("list[tuple[int, str]]", """
def f[T](t: T) -> list[T]: ...
expr = f((1, "hello"))
""");
}
@TestFor(issues="PY-57621")
public void testTupleAsGenericInTupleNarrows() {
var t = "tuple[Literal[1], Literal['hello']]";
doTest("tuple[list[tuple[int, str]], " + t + "]" + " | " + t, """
def f[T](t: T) -> tuple[list[T], T] | T: ...
expr = f((1, 'hello'))
""");
}
@TestFor(issues="PY-57621")
public void testTupleAsBareTypeVariableIsLiteral() {
doTest("tuple[Literal[1], Literal[\"hello\"]]", """
def f[T](t: T) -> T: ...
expr = f((1, "hello"))
""");
}
private void doTest(final String expectedType, final String text) {
myFixture.configureByText(PythonFileType.INSTANCE, text);
final PyExpression expr = myFixture.findElementByText("expr", PyExpression.class);
@@ -1002,7 +1002,7 @@ public class PyTypeTest extends PyTestCase {
// PY-9334
public void testIterateOverListOfNestedTuples() {
doTest("Literal['foo']",
doTest("str",
"""
def f():
for i, (expr, v) in [(0, ('foo', []))]:
@@ -1047,7 +1047,7 @@ public class PyTypeTest extends PyTestCase {
// PY-10967
public void testDefaultTupleParameterMember() {
doTest("Literal[1]",
doTest("int",
"""
def foo(xs=(1, 2)):
expr, foo = xs
@@ -1174,7 +1174,7 @@ public class PyTypeTest extends PyTestCase {
// PY-38928
public void testIterateListOfTuples() {
doTest(
"Literal['foo']",
"str",
"""
for ((_, expr)) in [(1, 'foo')]:
pass
@@ -2721,7 +2721,7 @@ public class PyTypeTest extends PyTestCase {
}
public void testUnpackingToNestedTargetsInSquareBracketsInForLoops() {
doTest("Literal[\"foo\"]",
doTest("str",
"""
xs = [(1, ("foo",))]
for [_, [expr]] in xs:
@@ -2730,7 +2730,7 @@ public class PyTypeTest extends PyTestCase {
}
public void testUnpackingToNestedTargetsInSquareBracketsInComprehensions() {
doTest("Literal[\"foo\"]",
doTest("str",
"""
xs = [(1, ("foo",))]
ys = [expr for [_, [expr]] in xs]
@@ -1,6 +1,7 @@
// Copyright 2000-2026 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package com.jetbrains.python.inspections;
import com.intellij.idea.TestFor;
import com.intellij.openapi.util.RecursionManager;
import com.intellij.openapi.util.StackOverflowPreventedException;
import com.jetbrains.python.fixtures.PyInspectionTestCase;
@@ -4371,5 +4372,20 @@ public class Py3TypeCheckerInspectionTest extends PyInspectionTestCase {
return set(self) # OK
""");
}
@TestFor(issues="PY-57621")
public void testTupleInGenericExplicitIsValid() {
doTestByText("""
from typing import Literal
class A[T]:
def __init__(self, t: T): ...
A[list[tuple[Literal[1]]]]([(1,)])
_: list[tuple[Literal[1]]] = [(1,)]
_: list[tuple[int]] = [(1,)]
""");
}
}
@@ -361,7 +361,7 @@ class PyTestFixtureResolvingTest : PyTestCase() {
}
fun testNamedParameterTypes() {
assertCorrectType(PARAMETRIZED_DIR, TEST_PARAMETER_TYPES, "Literal[9] | str")
assertCorrectType(PARAMETRIZED_DIR, TEST_PARAMETER_TYPES, INT_STR_UNION)
}
@TestFor(issues = ["PY-56268"])