diff --git a/python/src/com/jetbrains/python/psi/impl/PyReferenceExpressionImpl.java b/python/src/com/jetbrains/python/psi/impl/PyReferenceExpressionImpl.java index 3a6a5f97911f..350b15b8d603 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyReferenceExpressionImpl.java +++ b/python/src/com/jetbrains/python/psi/impl/PyReferenceExpressionImpl.java @@ -504,22 +504,23 @@ public class PyReferenceExpressionImpl extends PyElementImpl implements PyRefere final PyElement element = augAssignment != null ? augAssignment : anchor; try { final List defs = PyDefUseUtil.getLatestDefs(scopeOwner, name, element, true, false); - PyType type = null; - for (Instruction def : defs) { - final ReadWriteInstruction instruction = PyUtil.as(def, ReadWriteInstruction.class); - if (instruction != null) { - final Ref instructionType = instruction.getType(context, anchor); - if (instructionType != null) { - if (type != null) { - type = PyUnionType.union(type, instructionType.get()); - } - else { - type = instructionType.get(); - } + // null means empty set of possible types, Ref(null) means Any + final @Nullable Ref combinedType = StreamEx.of(defs) + .select(ReadWriteInstruction.class) + .map(instr -> instr.getType(context, anchor)) + // don't use foldLeft(BiFunction) here, as it doesn't support null results + .foldLeft(null, (accType, defType) -> { + if (defType == null) { + return accType; } - } - } - return type; + else if (accType == null) { + return defType; + } + else { + return Ref.create(PyUnionType.union(accType.get(), defType.get())); + } + }); + return Ref.deref(combinedType); } catch (PyDefUseUtil.InstructionNotFoundException ignored) { } diff --git a/python/testSrc/com/jetbrains/python/PyTypeTest.java b/python/testSrc/com/jetbrains/python/PyTypeTest.java index 44b7bcb278c9..3abba779fdfb 100644 --- a/python/testSrc/com/jetbrains/python/PyTypeTest.java +++ b/python/testSrc/com/jetbrains/python/PyTypeTest.java @@ -2001,6 +2001,17 @@ public class PyTypeTest extends PyTestCase { "expr = xs\n"); } + // PY-21175 + public void testAnyAddedByConditionalDefinition() { + doTest("Union[str, Any]", + "def f(x, y):\n" + + " if x:\n" + + " var = y\n" + + " else:\n" + + " var = 'foo'\n" + + " expr = var"); + } + private static List getTypeEvalContexts(@NotNull PyExpression element) { return ImmutableList.of(TypeEvalContext.codeAnalysis(element.getProject(), element.getContainingFile()).withTracing(), TypeEvalContext.userInitiated(element.getProject(), element.getContainingFile()).withTracing());