Skip to content

Commit e49725b

Browse files
committed
fix argo conditionals issue with new argo
1 parent 87a1df2 commit e49725b

3 files changed

Lines changed: 59 additions & 7 deletions

File tree

‎metaflow/plugins/argo/argo_workflows.py‎

Lines changed: 15 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1225,6 +1225,20 @@ def _skippable_input_steps_in_dag_order(self, node):
12251225
reverse=True,
12261226
)
12271227

1228+
def _input_path_ref(self, node_name):
1229+
sanitized = self._sanitize(node_name)
1230+
pred_node = self.graph[node_name]
1231+
if self._is_conditional_node(pred_node) or pred_node.type == "split-switch":
1232+
return (
1233+
"argo-{{workflow.name}}/%s/"
1234+
"{{=tasks['%s']?.outputs?.parameters['task-id'] ?? 'SKIPPED'}}"
1235+
% (node_name, sanitized)
1236+
)
1237+
return "argo-{{workflow.name}}/%s/{{tasks.%s.outputs.parameters.task-id}}" % (
1238+
node_name,
1239+
sanitized,
1240+
)
1241+
12281242
def _is_recursive_node(self, node):
12291243
return node.name in self.recursive_nodes
12301244

@@ -1378,11 +1392,7 @@ def _visit(
13781392
parameters = [
13791393
Parameter("input-paths").value(
13801394
compress_list(
1381-
[
1382-
"argo-{{workflow.name}}/%s/{{tasks.%s.outputs.parameters.task-id}}"
1383-
% (n, self._sanitize(n))
1384-
for n in node.in_funcs
1385-
],
1395+
[self._input_path_ref(n) for n in node.in_funcs],
13861396
# NOTE: We set zlibmin to infinite because zlib compression for the Argo input-paths breaks template value substitution.
13871397
zlibmin=inf,
13881398
)

‎metaflow/plugins/argo/conditional_input_paths.py‎

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -25,7 +25,9 @@ def generate_input_paths(input_paths, skippable_steps):
2525
# strip these out of the list.
2626

2727
# all pathspecs of leading steps that executed.
28-
trimmed = [path for path in paths if not "{{" in path]
28+
trimmed = [
29+
path for path in paths if "{{" not in path and not path.endswith("/SKIPPED")
30+
]
2931

3032
skippable_steps = [step for step in skippable_steps if step]
3133
skippable_step_set = set(skippable_steps)

‎test/unit/test_argo_conditional_input_paths.py‎

Lines changed: 41 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -75,6 +75,10 @@ def _unresolved_task_path(step_name):
7575
)
7676

7777

78+
def _skipped_task_path(step_name):
79+
return "%s/%s/SKIPPED" % (RUN_ID, step_name)
80+
81+
7882
def test_chain_skip_fallback_uses_latest_executed_split_switch(chain_skip_argo):
7983
node = chain_skip_argo.graph["end"]
8084
assert node.in_funcs == ["start", "step2", "step3"]
@@ -91,6 +95,20 @@ def test_chain_skip_fallback_uses_latest_executed_split_switch(chain_skip_argo):
9195
assert _decode_input_paths(result) == [_task_path("step2")]
9296

9397

98+
def test_chain_skip_with_skipped_sentinel(chain_skip_argo):
99+
"""Same as above but with SKIPPED sentinel (Argo v3.7.11+ behavior)."""
100+
node = chain_skip_argo.graph["end"]
101+
skippable_steps = chain_skip_argo._skippable_input_steps_in_dag_order(node)
102+
103+
input_paths = _encode_input_paths(
104+
[_task_path("start"), _task_path("step2"), _skipped_task_path("step3")]
105+
)
106+
107+
result = generate_input_paths(input_paths, skippable_steps)
108+
109+
assert _decode_input_paths(result) == [_task_path("step2")]
110+
111+
94112
@pytest.mark.parametrize(
95113
"paths, skippable_steps, expected",
96114
[
@@ -105,8 +123,30 @@ def test_chain_skip_fallback_uses_latest_executed_split_switch(chain_skip_argo):
105123
[_task_path("branch")],
106124
),
107125
([_task_path("step"), _task_path("step2")], ["step"], [_task_path("step2")]),
126+
(
127+
[_task_path("start"), _skipped_task_path("branch")],
128+
["start"],
129+
[_task_path("start")],
130+
),
131+
(
132+
[_task_path("start"), _task_path("step2"), _skipped_task_path("step3")],
133+
["step2", "start"],
134+
[_task_path("step2")],
135+
),
136+
(
137+
[_skipped_task_path("left"), _task_path("right")],
138+
[],
139+
[_task_path("right")],
140+
),
141+
],
142+
ids=[
143+
"normal_join",
144+
"non_skippable_executed",
145+
"exact_step_name",
146+
"skipped_non_skippable",
147+
"skipped_with_skippable_fallback",
148+
"skipped_no_skippable_steps",
108149
],
109-
ids=["normal_join", "non_skippable_executed", "exact_step_name"],
110150
)
111151
def test_generate_input_paths_filters_by_exact_step_name(
112152
paths, skippable_steps, expected

0 commit comments

Comments
 (0)