From 19900e42da7871e54d8ea45bd523741178cd3476 Mon Sep 17 00:00:00 2001 From: 1fanwang <1fannnw@gmail.com> Date: Sat, 19 Sep 2026 11:08:07 -0700 Subject: [PATCH] fix: resolve indexed promises in local conditions Signed-off-by: 1fanwang <1fannnw@gmail.com> --- flytekit/core/promise.py | 6 +++- tests/flytekit/unit/core/test_conditions.py | 38 +++++++++++++++++++++ 2 files changed, 43 insertions(+), 1 deletion(-) diff --git a/flytekit/core/promise.py b/flytekit/core/promise.py index 712f0d25ca..9cc888f647 100644 --- a/flytekit/core/promise.py +++ b/flytekit/core/promise.py @@ -286,12 +286,14 @@ class ComparisonExpression(object): and operator can be any comparison expression like <, >, <=, >=, ==, != """ - def __init__(self, lhs: Union["Promise", Any], op: ComparisonOps, rhs: Union["Promise", Any]): + def __init__(self, lhs: Union["Promise", Any], op: ComparisonOps, rhs: Union["Promise", Any]) -> None: self._op = op self._lhs = None self._rhs = None if isinstance(lhs, Promise): + if lhs.is_ready and lhs.attr_path: + lhs = run_sync(coro_func=resolve_attr_path_in_promise, p=lhs.deepcopy()) self._lhs = lhs if lhs.is_ready: if lhs.val.scalar is None or lhs.val.scalar.primitive is None: @@ -304,6 +306,8 @@ def __init__(self, lhs: Union["Promise", Any], op: ComparisonOps, rhs: Union["Pr else: raise ValueError("Only primitive values can be used in comparison") if isinstance(rhs, Promise): + if rhs.is_ready and rhs.attr_path: + rhs = run_sync(coro_func=resolve_attr_path_in_promise, p=rhs.deepcopy()) self._rhs = rhs if rhs.is_ready: if rhs.val.scalar is None or rhs.val.scalar.primitive is None: diff --git a/tests/flytekit/unit/core/test_conditions.py b/tests/flytekit/unit/core/test_conditions.py index 2192590387..3da5e85367 100644 --- a/tests/flytekit/unit/core/test_conditions.py +++ b/tests/flytekit/unit/core/test_conditions.py @@ -215,6 +215,44 @@ def decompose_unary() -> int: return conditional("test").if_(result.is_none()).then(success()).else_().then(failed()) +@pytest.mark.parametrize("value", [None, True, False]) +def test_condition_is_none_on_indexed_promise(value: typing.Optional[bool]) -> None: + @task + def status(value: typing.Optional[bool]) -> typing.List[typing.Optional[bool]]: + return [value, True] + + @workflow + def check(value: typing.Optional[bool]) -> float: + result = status(value=value)[0] + return conditional("test").if_(result.is_none()).then(square(n=3.0)).else_().then(double(n=3.0)) + + assert check(value=value) == (9.0 if value is None else 6.0) + + wf_spec = get_serializable( + entity_mapping=OrderedDict(), + settings=serialization_settings, + entity=check, + ) + branch = wf_spec.template.nodes[1] + assert branch.inputs[0].binding.promise.attr_path == [0] + assert branch.branch_node.if_else.case.condition.comparison.right_value.scalar.none_type is not None + + +@pytest.mark.parametrize(("value", "expected"), [(4, 9.0), (5, 6.0), (6, 6.0)]) +def test_condition_with_indexed_rhs(value: int, expected: float) -> None: + @task + def values(value: int) -> typing.List[int]: + return [value] + + @workflow + def check(value: int) -> float: + left = five() + right = values(value=value)[0] + return conditional("test").if_(left > right).then(square(n=3.0)).else_().then(double(n=3.0)) + + assert check(value=value) == expected + + def test_subworkflow_condition_serialization(): """Test that subworkflows are correctly extracted from serialized workflows with condiationals."""