PY-57621 type inference: infer tuples with literal types

(cherry picked from commit d3f6711e4d187fac8b426eb1d5ed75def4932b3b)

GitOrigin-RevId: cc327f405e9a133425bcfbb091e22f1e267b6383
This commit is contained in:
Morgan Bartholomew
2026-03-25 07:22:50 +00:00
committed by intellij-monorepo-bot
parent 90bbd2ac21
commit 4d95ecde4e
25 changed files with 121 additions and 52 deletions
@@ -244,6 +244,13 @@ public final class PyStdlibTypeProvider extends PyTypeProviderBase {
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))
);
}
return getEnumAttributeInfo(enumClass, type, context);
}
else {
@@ -10,6 +10,7 @@ import com.jetbrains.python.psi.PyTupleExpression
import com.jetbrains.python.psi.types.PyTupleType
import com.jetbrains.python.psi.types.PyType
import com.jetbrains.python.psi.types.TypeEvalContext
import com.jetbrains.python.psi.types.getLiteralType
class PyTupleExpressionImpl(astNode: ASTNode) : PySequenceExpressionImpl(astNode), PyTupleExpression {
override fun acceptPyVisitor(pyVisitor: PyElementVisitor) {
@@ -19,7 +20,7 @@ class PyTupleExpressionImpl(astNode: ASTNode) : PySequenceExpressionImpl(astNode
override fun getType(context: TypeEvalContext, key: TypeEvalContext.Key): PyType? {
return PyTupleType.create(
this,
elements.map { context.getType(it) }
elements.map { it.getLiteralType(context) ?: context.getType(it) }
)
}
@@ -257,7 +257,7 @@ class PyLiteralType private constructor(cls: PyClass, val expression: PyExpressi
return PyLiteralStringType.create(expression)
}
}
return getLiteralType(expression, context)
return expression.getLiteralType(context)
}
private fun literalType(expression: PyExpression, context: TypeEvalContext, index: Boolean): PyLiteralType? {
@@ -436,7 +436,12 @@ object PyTypeChecker {
context
)
}
context.mySubstitutions.putTypeVarTuple(expected as PyTypeVarTupleType, actual, KeyImpl)
val normalizedActual =
if (actual is PyUnpackedTupleType)
// TODO: consider how widening should work with more complex types like: `tuple[Sequence[Literal[1]]`
PyUnpackedTupleTypeImpl(actual.elementTypes.map { PyLiteralType.upcastLiteralToClass(it) }, actual.isUnbound)
else actual
context.mySubstitutions.putTypeVarTuple(expected as PyTypeVarTupleType, normalizedActual, KeyImpl)
}
return true
}
@@ -19,6 +19,7 @@ import com.intellij.openapi.util.Key
import com.intellij.openapi.util.Ref
import com.intellij.openapi.util.UserDataHolder
import com.intellij.psi.PsiElement
import com.jetbrains.python.psi.PyExpression
import com.jetbrains.python.psi.PyPsiFacade
import com.jetbrains.python.psi.impl.PyBuiltinCache
import com.jetbrains.python.psi.types.PyRecursiveTypeVisitor.PyTypeTraverser
@@ -368,3 +369,7 @@ val PyType?.isUnknown: Boolean
PyAnyType.validate(this)
return if (PyAnyType.isEnabled) this is PyAnyType.Unknown else this == null
}
@ApiStatus.Internal
fun PyExpression.getLiteralType(context: TypeEvalContext): PyType? =
PyLiteralType.getLiteralType(this, context)
@@ -1 +1 @@
<warning descr="'(int, int)' object is not callable">(1,2)()</warning>
<warning descr="'(Literal[1], Literal[2])' object is not callable">(1,2)()</warning>
@@ -100,9 +100,9 @@ print '%d, %d, %d, %d' % <warning descr="Too few arguments for format string">my
# PY-12801
print '%d %s' % ((42,) + ('spam',))
print '%d %s' % (<warning descr="Unexpected type (str, str)">('ham',) + ('spam',)</warning>)
print '%d %s' % (<warning descr="Too few arguments for format string"><warning descr="Unexpected type (int)">(42,) + ()</warning></warning>)
print '%d' % (<warning descr="Too many arguments for format string"><warning descr="Unexpected type (int, str)">(42,) + ('spam',)</warning></warning>)
print '%d %s' % (<warning descr="Unexpected type (Literal['ham'], Literal['spam'])">('ham',) + ('spam',)</warning>)
print '%d %s' % (<warning descr="Too few arguments for format string"><warning descr="Unexpected type (Literal[42])">(42,) + ()</warning></warning>)
print '%d' % (<warning descr="Too many arguments for format string"><warning descr="Unexpected type (Literal[42], Literal['spam'])">(42,) + ('spam',)</warning></warning>)
# PY-11274
import collections
@@ -1,2 +1,2 @@
args = ('foo', 'bar')
s = '%d %d' % <warning descr="Unexpected type (str, str)">args</warning>
s = '%d %d' % <warning descr="Unexpected type (Literal['foo'], Literal['bar'])">args</warning>
@@ -2,4 +2,4 @@ argument_pattern = re.compile(r'(%s)\s*(\(\s*(%s)\s*\)\s*)?$'
% ((states.Inliner.simplename,) * 2))
t, num = ('foo',), 2
res = '%d %d' % (<warning descr="Unexpected type (str, str)">t * num</warning>)
res = '%d %d' % (<warning descr="Unexpected type (Literal['foo'], Literal['foo'])">t * num</warning>)
@@ -1,6 +1,6 @@
print('a' < 'b' < 'c' < 'd')
print(('a' < 'b') < <warning descr="Expected type 'int', got 'str' instead">'c'</warning>)
print((1, 1) < (1, 2) < (1, 3) < (1, 4))
print((1, 1) < (1, 2) < <warning descr="Expected type 'Tuple[Literal[1, 2], ...]' (matched generic type 'Tuple[_T_co, ...]'), got 'Tuple[Literal[1], Literal[3]]' instead">(1, 3)</warning> < <warning descr="Expected type 'Tuple[Literal[1, 3], ...]' (matched generic type 'Tuple[_T_co, ...]'), got 'Tuple[Literal[1], Literal[4]]' instead">(1, 4)</warning>)
print(((1, 1) < (1, 2)) < <warning descr="Expected type 'int', got 'Tuple[int, int]' instead">(1, 3)</warning>)
print(1.0 < 4.5 < 9.3 < 10.0)
print((1.0 < 4.5) < 9.3)
@@ -9,7 +9,7 @@ int_and_bool = (42, True)
expects_many_ints(int_and_bool)
int_and_str = (42, 'foo')
expects_many_ints(<warning descr="Expected type 'tuple[int, ...]', got 'tuple[int, str]' instead">int_and_str</warning>)
expects_many_ints(<warning descr="Expected type 'tuple[int, ...]', got 'tuple[Literal[42], Literal['foo']]' instead">int_and_str</warning>)
booleans = (True, False) # type: Tuple[bool, ...]
expects_many_ints(booleans)
@@ -29,7 +29,7 @@ d[frozenset([1, 2])] = 0
d[object()] = 0
d[(1, (2, 3))] = 0
d[<error descr="Cannot use unhashable type '(int, (int, list))' as a dict key">(1, (2, []))</error>] = 0
d[<error descr="Cannot use unhashable type '(Literal[1], (Literal[2], list))' as a dict key">(1, (2, []))</error>] = 0
unhashable_union: int | list = 5
d[<error descr="Cannot use unhashable type 'int | list' as a dict key">unhashable_union</error>] = 0
@@ -1 +1,3 @@
var: [tuple[int, int]] = (1, 2)
from typing import Literal
var: [tuple[Literal[1], Literal[2]]] = (1, 2)
@@ -1,4 +1,7 @@
from typing import Literal
def func():
var: [str]
var: [Literal['spam']]
var, _ = 'spam', 42
var
@@ -1 +1,3 @@
var: [tuple[int, str, None]] = (1, 'foo', None)
from typing import Literal
var: [tuple[Literal[1], Literal['foo'], None]] = (1, 'foo', None)
@@ -1,3 +1,6 @@
from typing import Literal
def func():
((var, _), _) = ('foo', 1), 2 # type: (([str], [int]), [int])
((var, _), _) = ('foo', 1), 2 # type: (([Literal['foo']], [Literal[1]]), [Literal[2]])
var
@@ -1,3 +1,6 @@
from typing import Literal
def func():
var, _ = 'spam', 42 # type: ([str], [int])
var, _ = 'spam', 42 # type: ([Literal['spam']], [Literal[42]])
var
@@ -1 +1 @@
<html><body><div class="bottom"><icon src="AllIcons.Nodes.Package"/>&nbsp;<code><a href="psi_element://#module#TupleTypeIsRenderedLowercased">TupleTypeIsRenderedLowercased</a></code></div><div class="definition"><pre><span style="color:#000000;">items</span><span style="">: </span><span style="color:#000000;"><span style="color:#000080;">tuple</span><span style="">[</span><span style="color:#000080;"><a href="psi_element://#typename#int">int</a></span><span style="">, </span><span style="color:#000080;"><a href="psi_element://#typename#str">str</a></span><span style="">]</span></span><span style=""> = </span><span style="">(</span><span style="color:#0000ff;">42</span><span style="">,&#32;</span><span style="color:#008000;font-weight:bold;">'foo'</span><span style="">)</span></pre></div></body></html>
<html><body><div class="bottom"><icon src="AllIcons.Nodes.Package"/>&nbsp;<code><a href="psi_element://#module#TupleTypeIsRenderedLowercased">TupleTypeIsRenderedLowercased</a></code></div><div class="definition"><pre><span style="color:#000000;">items</span><span style="">: </span><span style="color:#000000;"><span style="color:#000080;">tuple</span><span style="">[</span>Literal[42]<span style="">, </span>Literal['foo']<span style="">]</span></span><span style=""> = </span><span style="">(</span><span style="color:#0000ff;">42</span><span style="">,&#32;</span><span style="color:#008000;font-weight:bold;">'foo'</span><span style="">)</span></pre></div></body></html>
@@ -3646,7 +3646,7 @@ public class Py3TypeTest extends PyTestCase {
// PY-64474
public void testTupleElementAccessedWithNegativeIndex() {
doTest("bool",
doTest("Literal[True]",
"""
xs = (1, True, "foo")
expr = xs[-2]
@@ -5151,6 +5151,13 @@ public class Py3TypeTest extends PyTestCase {
""");
}
@TestFor(issues = "PY-57621")
public void testTupleWithLiteralValues() {
doTest("tuple[Literal[1]]", """
expr = (1,)
""");
}
// PY-87575
public void testIterDefinedInMetaclass() {
doTest("set[int]", """
@@ -1,4 +1,4 @@
// Copyright 2000-2025 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
// 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;
import com.jetbrains.python.documentation.PythonDocumentationProvider;
@@ -32,7 +32,7 @@ public final class PyTypeConversionTest extends PyTestCase {
}
public void testTupleToTypingIterable() {
doTest("typing.Iterable", "Iterable[int | str]", """
doTest("typing.Iterable", "Iterable[Literal[1, \"foo\"]]", """
expr = (1, "foo")
""");
}
@@ -1,4 +1,4 @@
// Copyright 2000-2021 JetBrains s.r.o. Use of this source code is governed by the Apache 2.0 license that can be found in the LICENSE file.
// 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;
import com.google.common.collect.ImmutableList;
@@ -52,7 +52,7 @@ public class PyTypeTest extends PyTestCase {
}
public void testTupleType() {
doTest("str",
doTest("Literal['a']",
"t = ('a', 2)\n" +
"expr = t[0]");
doTest("List[bool]",
@@ -68,7 +68,7 @@ public class PyTypeTest extends PyTestCase {
}
public void testTupleAssignmentType() {
doTest("str",
doTest("Literal['a']",
"t = ('a', 2)\n" +
"(expr, q) = t");
}
@@ -106,7 +106,7 @@ public class PyTypeTest extends PyTestCase {
}
public void testUnionOfTuples() {
doTest("Union[Tuple[int, str], Tuple[str, int]]",
doTest("Union[Tuple[Literal[1], Literal['a']], Tuple[Literal['a'], Literal[1]]]",
"""
def x(b):
if b:
@@ -298,7 +298,7 @@ public class PyTypeTest extends PyTestCase {
}
public void testIsInstance2() {
doTest("str",
doTest("Literal[\"\"]",
"""
x = ""
if isinstance(x, (1, "")):
@@ -1002,7 +1002,7 @@ public class PyTypeTest extends PyTestCase {
// PY-9334
public void testIterateOverListOfNestedTuples() {
doTest("str",
doTest("Literal['foo']",
"""
def f():
for i, (expr, v) in [(0, ('foo', []))]:
@@ -1047,7 +1047,7 @@ public class PyTypeTest extends PyTestCase {
// PY-10967
public void testDefaultTupleParameterMember() {
doTest("int",
doTest("Literal[1]",
"""
def foo(xs=(1, 2)):
expr, foo = xs
@@ -1071,7 +1071,7 @@ public class PyTypeTest extends PyTestCase {
}
public void testTupleFromTuple() {
doTest("Tuple[str, int, int]",
doTest("Tuple[Literal['1'], Literal[2], Literal[3]]",
"expr = tuple(('1', 2, 3))");
}
@@ -1128,7 +1128,7 @@ public class PyTypeTest extends PyTestCase {
}
public void testTupleIterationType() {
doTest("Union[int, str]",
doTest("Literal[1, 'a']",
"""
xs = (1, 'a')
for expr in xs:
@@ -1138,35 +1138,35 @@ public class PyTypeTest extends PyTestCase {
// PY-12801
public void testTupleConcatenation() {
doTest("Tuple[int, bool, str]",
doTest("Tuple[Literal[1], Literal[True], Literal['spam']]",
"expr = (1,) + (True, 'spam') + ()");
}
public void testTupleMultiplication() {
doTest("Tuple[int, bool, int, bool]",
doTest("Tuple[Literal[1], Literal[False], Literal[1], Literal[False]]",
"expr = (1, False) * 2");
}
public void testTupleDestructuring() {
doTest("str",
doTest("Literal['val']",
"_, expr = (1, 'val') ");
}
public void testParensTupleDestructuring() {
doTest("str",
doTest("Literal['val']",
"(_, expr) = (1, 'val') ");
}
// PY-19825
public void testSubTupleDestructuring() {
doTest("str",
doTest("Literal['val']",
"(a, (_, expr)) = (1, (2,'val')) ");
}
// PY-19825
public void testSubTupleIndirectDestructuring() {
doTest("str",
doTest("Literal['val']",
"xs = (2,'val')\n" +
"(a, (_, expr)) = (1, xs) ");
}
@@ -1174,7 +1174,7 @@ public class PyTypeTest extends PyTestCase {
// PY-38928
public void testIterateListOfTuples() {
doTest(
"str",
"Literal['foo']",
"""
for ((_, expr)) in [(1, 'foo')]:
pass
@@ -1670,9 +1670,11 @@ public class PyTypeTest extends PyTestCase {
}
public void testHeterogeneousTupleLiteral() {
doTest("Tuple[str, int, int]", "expr = ('1', 1, 1)");
doTest("Tuple[Literal['1'], Literal[1], Literal[1]]", "expr = ('1', 1, 1)");
doTest("Tuple[str, int, int, int, int, int, int, int, int, int, int]", "expr = ('1', 1, 1, 1, 1, 1, 1, 1, 1, 1, 1)");
doTest(
"Tuple[Literal['1'], Literal[1], Literal[1], Literal[1], Literal[1], Literal[1], Literal[1], Literal[1], Literal[1], Literal[1], Literal[1]]",
"expr = ('1', 1, 1, 1, 1, 1, 1, 1, 1, 1, 1)");
}
// PY-20818
@@ -2712,14 +2714,14 @@ public class PyTypeTest extends PyTestCase {
}
public void testUnpackingToNestedTargetsInSquareBracketsInAssignments() {
doTest("int",
doTest("Literal[42]",
"""
[_, [[expr], _]] = "foo", ((42,), "bar")
""");
}
public void testUnpackingToNestedTargetsInSquareBracketsInForLoops() {
doTest("str",
doTest("Literal[\"foo\"]",
"""
xs = [(1, ("foo",))]
for [_, [expr]] in xs:
@@ -2728,7 +2730,7 @@ public class PyTypeTest extends PyTestCase {
}
public void testUnpackingToNestedTargetsInSquareBracketsInComprehensions() {
doTest("str",
doTest("Literal[\"foo\"]",
"""
xs = [(1, ("foo",))]
ys = [expr for [_, [expr]] in xs]
@@ -6158,6 +6158,19 @@ public class PyTypingTest extends PyTestCase {
""");
}
@TestFor(issues="PY-57621")
public void testEnumTuple() {
doTest("tuple[int, str]", """
from enum import Enum
class Color(Enum):
RED = 1, "red"
BLUE = 2, "blue"
expr = Color.BLUE.value
""");
}
// PY-76149
public void testDataclassTransformConstructorSignatureWithFieldsAnnotatedWithDescriptor() {
doTestExpressionUnderCaret("(id: int, name: str) -> MyClass", """
@@ -1,4 +1,4 @@
// Copyright 2000-2017 JetBrains s.r.o. Use of this source code is governed by the Apache 2.0 license that can be found in the LICENSE file.
// 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.openapi.util.RecursionManager;
@@ -1874,8 +1874,8 @@ public class Py3TypeCheckerInspectionTest extends PyInspectionTestCase {
foo(1, bar, args=(0, 'foo'))
foo(1, baz, args=(0, 'foo', 1.0, False))
foo(1, bar, <warning descr="Expected type 'tuple[int, str]' (matched generic type 'tuple[*Ts]'), got 'tuple[str, int]' instead">args=('foo', 0)</warning>)
foo(1, baz, <warning descr="Expected type 'tuple[int, str, float, bool]' (matched generic type 'tuple[*Ts]'), got 'tuple[str, int, float, bool]' instead">args=('foo', 0, 1.0, False)</warning>)
foo(1, bar, <warning descr="Expected type 'tuple[int, str]' (matched generic type 'tuple[*Ts]'), got 'tuple[Literal['foo'], Literal[0]]' instead">args=('foo', 0)</warning>)
foo(1, baz, <warning descr="Expected type 'tuple[int, str, float, bool]' (matched generic type 'tuple[*Ts]'), got 'tuple[Literal['foo'], Literal[0], float, Literal[False]]' instead">args=('foo', 0, 1.0, False)</warning>)
""");
}
@@ -2044,12 +2044,28 @@ public class Py3TypeCheckerInspectionTest extends PyInspectionTestCase {
def foo(*args: Tuple[*Ts]): ...
foo((0,), (1,))
foo((0,), <warning descr="Expected type 'tuple[int]' (matched generic type 'tuple[*Ts]'), got 'tuple[int, int]' instead">(1, 2)</warning>)
foo((0,), <warning descr="Expected type 'tuple[int]' (matched generic type 'tuple[*Ts]'), got 'tuple[Literal[1], Literal[2]]' instead">(1, 2)</warning>)
# Should fail according to https://typing.python.org/en/latest/spec/generics.html#type-variable-tuple-equality
foo((0,), <warning descr="Expected type 'tuple[int]' (matched generic type 'tuple[*Ts]'), got 'tuple[str]' instead">('1',)</warning>)
foo((0,), <warning descr="Expected type 'tuple[int]' (matched generic type 'tuple[*Ts]'), got 'tuple[Literal['1']]' instead">('1',)</warning>)
""");
}
public void testTypeVarTupleWidening() {
fixme("widen more literal types in type var tuples", AssertionError.class, "Expected type 'tuple[tuple[Literal[0]]]'", () -> {
doTestByText("""
from typing import Literal, Sequence
def foo[*Ts](*args: tuple[*Ts]): ...
# nested tuples
foo(((0,),), ((1,),))
def main(ones: Sequence[Literal[1]], twos: Sequence[Literal[2]]):
# should this widen to `Sequence[int]` or should it show an error?
foo((ones,), (twos,))
""");
});
}
// PY-53105
public void testVariadicGenericStarArgsOfVariadicGenericPrefixSuffix() {
doTestByText("""
@@ -1,4 +1,4 @@
// Copyright 2000-2025 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
// 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.jetbrains.python.fixtures.PyInspectionTestCase;
@@ -141,7 +141,7 @@ def f(a):
class D:
def __init__(self):
self.x = 0
__match_args__ = (<warning descr="Expected type 'tuple[str, ...]', got 'tuple[str, int]' instead">"x", 1</warning>)
__match_args__ = (<warning descr="Expected type 'tuple[str, ...]', got 'tuple[Literal[\\"x\\"], Literal[1]]' instead">"x", 1</warning>)
""");
}
@@ -278,7 +278,7 @@ class D:
public void testMatchArgsInvalidTupleOfInts() {
doTestByText("""
class D:
__match_args__ = (<warning descr="Expected type 'tuple[str, ...]', got 'tuple[int, int, int]' instead">1, 2, 3</warning>)
__match_args__ = (<warning descr="Expected type 'tuple[str, ...]', got 'tuple[Literal[1], Literal[2], Literal[3]]' instead">1, 2, 3</warning>)
""");
}
@@ -1,4 +1,4 @@
// Copyright 2000-2023 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
// 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.testing
import com.intellij.codeInsight.navigation.actions.GotoTypeDeclarationAction
@@ -361,7 +361,7 @@ class PyTestFixtureResolvingTest : PyTestCase() {
}
fun testNamedParameterTypes() {
assertCorrectType(PARAMETRIZED_DIR, TEST_PARAMETER_TYPES, INT_STR_UNION)
assertCorrectType(PARAMETRIZED_DIR, TEST_PARAMETER_TYPES, "Literal[9] | str")
}
@TestFor(issues = ["PY-56268"])