diff --git a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyLiteralType.kt b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyLiteralType.kt index d01bd344396b..71e4ad12c793 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyLiteralType.kt +++ b/python/python-psi-impl/src/com/jetbrains/python/psi/types/PyLiteralType.kt @@ -43,6 +43,7 @@ import com.jetbrains.python.psi.impl.PyPsiFacadeImpl import com.jetbrains.python.psi.resolve.PyResolveContext import com.jetbrains.python.psi.resolve.PyResolveUtil import com.jetbrains.python.psi.stubs.PyLiteralKind +import com.jetbrains.python.psi.types.PyTypeUtil.asUnionSequence import com.jetbrains.python.psi.types.PyTypeUtil.toStream import org.jetbrains.annotations.ApiStatus import java.math.BigInteger @@ -276,11 +277,13 @@ class PyLiteralType private constructor( } private fun promoteDictLiteral(expectedType: PyType?, dictLiteral: PyDictLiteralExpression): PyType? { - if (expectedType is PyTypedDictType) { - val typeCheckingResult = PyTypedDictType.TypeCheckingResult() - PyTypedDictType.checkExpression(expectedType, dictLiteral, context, typeCheckingResult) - if (!typeCheckingResult.hasErrors) { - return expectedType + for (candidate in expectedType.asUnionSequence()) { + if (candidate is PyTypedDictType) { + val typeCheckingResult = PyTypedDictType.TypeCheckingResult() + PyTypedDictType.checkExpression(candidate, dictLiteral, context, typeCheckingResult) + if (!typeCheckingResult.hasErrors) { + return candidate + } } } return promoteDictLiteralOrDictComprehension(expectedType, dictLiteral, dictLiteral.elements.asList()) diff --git a/python/testSrc/com/jetbrains/python/types/PyTypedDictTypeTest.kt b/python/testSrc/com/jetbrains/python/types/PyTypedDictTypeTest.kt index f386565b940a..d9b197916b8b 100644 --- a/python/testSrc/com/jetbrains/python/types/PyTypedDictTypeTest.kt +++ b/python/testSrc/com/jetbrains/python/types/PyTypedDictTypeTest.kt @@ -528,6 +528,39 @@ class PyTypedDictTypeTest : PyCodeInsightTestCase() { # ^^^^^^^^^^^^^^^^^^^^^^^^^^ WARNING TypedDict 'Movie' has missing key: 'name' """) + @Test + @TestFor(issues = ["PY-88391"]) + fun `dict literal assignable to optional TypedDict`() = test(""" + from typing import TypedDict + + class Address(TypedDict): + street: str + + a: Address | None = {"street": "Pine"} + b: Address | None = {"color": "red"} + # ^^^^^^^^^^^^^^^^ WARNING Expected type 'Address | None', got 'dict[Literal["color"], Literal["red"]]' instead + """) + + @Test + @TestFor(issues = ["PY-88391"]) + fun `dict literal assignable to one of several TypedDicts`() = test(""" + from typing import TypedDict + + class A(TypedDict): + a: str + + class B(TypedDict): + b: int + + class C(TypedDict): + c: int + + first: A | B | C = {"a": "x"} + last: A | B | C = {"c": 1} + none: A | B | C = {"z": 1} + # ^^^^^^^^ WARNING Expected type 'A | B | C', got 'dict[Literal["z"], Literal[1]]' instead + """) + @Test @TestFor(issues = ["PY-38873"]) fun `value access through list field`() = test(TestOptions(enablePyAnyType = false), """