Skip to content

Commit b709e32

Browse files
committed
refactor(tpu-observability): wrap TPU observability DAGs with TaskGroupWithTimeout
Refactors five TPU observability DAGs to use TaskGroupWithTimeout for proper timeout management during pre-test setup, test execution, and post-test teardown. This aligns their implementation pattern with node_pool_status.py. Specifically, this change: - Imports TaskGroupWithTimeout in each DAG. - Defines PRE_TEST_TIMEOUT, TEST_TIMEOUT, and POST_TEST_TIMEOUT at the module level. - Groups tasks into pre_test, test, and post_test TaskGroupWithTimeout blocks. - Chains tasks sequentially within each group and updates the overall DAG chain. Modified DAGs: - dags/tpu_observability/jobset_ttr_drain_restart.py - dags/tpu_observability/jobset_ttr_node_pool_resize.py - dags/tpu_observability/jobset_ttr_pod_delete.py - dags/tpu_observability/jobset_ttr_rollback.py - dags/tpu_observability/jobset_uptime_validation.py
1 parent c4dbf83 commit b709e32

5 files changed

Lines changed: 488 additions & 324 deletions

File tree

dags/tpu_observability/jobset_ttr_drain_restart.py

Lines changed: 102 additions & 67 deletions
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,11 @@
2121
from airflow.models.baseoperator import chain
2222
from airflow.utils.task_group import TaskGroup
2323
from airflow.utils.trigger_rule import TriggerRule
24+
from airflow.utils.task_group import TaskGroup
25+
from dags.common.task_group_with_timeout import TaskGroupWithTimeout
26+
27+
28+
from airflow.decorators import task
2429

2530
from dags import composer_env
2631
from dags.common.scheduling_helper.scheduling_helper import (
@@ -42,6 +47,10 @@
4247
DAGRUN_TIMEOUT = get_dag_timeout(DAG_ID)
4348
SCHEDULE = SchedulingHelper.arrange_schedule_time(DAG_ID)
4449

50+
PRE_TEST_TIMEOUT = datetime.timedelta(minutes=10)
51+
POST_TEST_TIMEOUT = datetime.timedelta(minutes=10)
52+
TEST_TIMEOUT = DAGRUN_TIMEOUT - PRE_TEST_TIMEOUT - POST_TEST_TIMEOUT
53+
4554

4655
@task
4756
def check_nodes_number(
@@ -140,75 +149,101 @@ def check_nodes_number(
140149
selector = jobset.generate_node_pool_selector(DAG_ID)
141150
jobset_name = jobset.generate_jobset_name(jobset_config.dag_id_prefix)
142151

143-
create_node_pool = node_pool.create.override(task_id="create_node_pool")(
144-
node_pool=cluster_info,
145-
node_pool_selector=selector,
146-
)
147-
148-
startup = jobset.create_jobset_startup_tasks(
149-
node_pool=cluster_info,
150-
jobset_config=jobset_config,
151-
jobset_name=jobset_name,
152-
node_pool_selector=selector,
153-
workload_type=Workload.JAX_TPU_BENCHMARK,
154-
)
155-
156-
select_node = node_pool.draw_random_node.override(task_id="select_node")(
157-
node_pool=cluster_info
158-
)
159-
160-
drained_node = node_pool.operate_node.override(task_id="drained_node")(
161-
node_pool=cluster_info,
162-
operation=NodeOperationSpec.Drain(),
163-
node_name=select_node,
164-
)
165-
166-
check_nodes_number = check_nodes_number.override(
167-
task_id="check_nodes_number"
168-
)(
169-
pool=cluster_info,
170-
drained_node_number=1,
171-
)
172-
173-
uncordon_node = node_pool.operate_node.override(task_id="uncordon_node")(
174-
node_pool=cluster_info,
175-
operation=NodeOperationSpec.Uncordon(),
176-
node_name=select_node,
177-
)
178-
179-
wait_for_metric_upload = jobset.wait_for_jobset_ttr_to_be_found.override(
180-
task_id="wait_for_metric_upload"
181-
)(
182-
node_pool=cluster_info,
183-
jobset_name=jobset_name,
184-
)
185-
186-
cleanup_workload = jobset.end_workload.override(
187-
task_id="cleanup_workload", trigger_rule=TriggerRule.ALL_DONE
188-
)(
189-
node_pool=cluster_info,
190-
jobset_config=jobset_config,
191-
jobset_name=jobset_name,
192-
).as_teardown(
193-
setups=startup.jobset_start_time
194-
)
195-
196-
cleanup_node_pool = node_pool.delete.override(
197-
task_id="cleanup_node_pool", trigger_rule=TriggerRule.ALL_DONE
198-
)(node_pool=cluster_info).as_teardown(
199-
setups=create_node_pool,
200-
)
152+
with TaskGroupWithTimeout(
153+
group_id="pre_test",
154+
timeout=PRE_TEST_TIMEOUT,
155+
) as pre_test:
156+
create_node_pool = node_pool.create.override(
157+
task_id="create_node_pool"
158+
)(
159+
node_pool=cluster_info,
160+
node_pool_selector=selector,
161+
)
162+
163+
startup = jobset.create_jobset_startup_tasks(
164+
node_pool=cluster_info,
165+
jobset_config=jobset_config,
166+
jobset_name=jobset_name,
167+
node_pool_selector=selector,
168+
workload_type=Workload.JAX_TPU_BENCHMARK,
169+
)
170+
171+
with TaskGroupWithTimeout(
172+
group_id="test",
173+
timeout=TEST_TIMEOUT,
174+
) as test:
175+
select_node = node_pool.draw_random_node.override(
176+
task_id="select_node"
177+
)(node_pool=cluster_info)
178+
179+
drained_node = node_pool.operate_node.override(task_id="drained_node")(
180+
node_pool=cluster_info,
181+
operation=NodeOperationSpec.Drain(),
182+
node_name=select_node,
183+
)
184+
185+
check_nodes_number_task = check_nodes_number.override(
186+
task_id="check_nodes_number"
187+
)(
188+
pool=cluster_info,
189+
drained_node_number=1,
190+
)
191+
192+
uncordon_node = node_pool.operate_node.override(
193+
task_id="uncordon_node"
194+
)(
195+
node_pool=cluster_info,
196+
operation=NodeOperationSpec.Uncordon(),
197+
node_name=select_node,
198+
)
199+
200+
wait_for_metric_upload = (
201+
jobset.wait_for_jobset_ttr_to_be_found.override(
202+
task_id="wait_for_metric_upload"
203+
)(
204+
node_pool=cluster_info,
205+
jobset_name=jobset_name,
206+
)
207+
)
208+
209+
chain(
210+
select_node,
211+
drained_node,
212+
check_nodes_number_task,
213+
uncordon_node,
214+
wait_for_metric_upload,
215+
)
216+
217+
with TaskGroupWithTimeout(
218+
group_id="post_test",
219+
timeout=POST_TEST_TIMEOUT,
220+
is_teardown=True,
221+
) as post_test:
222+
cleanup_workload = jobset.end_workload.override(
223+
task_id="cleanup_workload", trigger_rule=TriggerRule.ALL_DONE
224+
)(
225+
node_pool=cluster_info,
226+
jobset_config=jobset_config,
227+
jobset_name=jobset_name,
228+
).as_teardown(
229+
setups=startup.jobset_start_time
230+
)
231+
232+
cleanup_node_pool = node_pool.delete.override(
233+
task_id="cleanup_node_pool", trigger_rule=TriggerRule.ALL_DONE
234+
)(node_pool=cluster_info).as_teardown(
235+
setups=create_node_pool,
236+
)
237+
238+
chain(
239+
cleanup_workload,
240+
cleanup_node_pool,
241+
)
201242

202243
chain(
203244
selector,
204245
jobset_name,
205-
create_node_pool,
206-
*startup.tasks,
207-
select_node,
208-
drained_node,
209-
check_nodes_number,
210-
uncordon_node,
211-
wait_for_metric_upload,
212-
cleanup_workload,
213-
cleanup_node_pool,
246+
pre_test,
247+
test,
248+
post_test,
214249
)

dags/tpu_observability/jobset_ttr_node_pool_resize.py

Lines changed: 98 additions & 68 deletions
Original file line numberDiff line numberDiff line change
@@ -34,12 +34,18 @@
3434
from dags.tpu_observability.utils import jobset_util as jobset
3535
from dags.tpu_observability.utils import node_pool_util as node_pool
3636
from dags.tpu_observability.utils.jobset_util import Workload
37+
from dags.common.task_group_with_timeout import TaskGroupWithTimeout
38+
3739

3840
DAG_ID = "jobset_ttr_node_pool_resize"
3941
DAGRUN_TIMEOUT = get_dag_timeout(DAG_ID)
4042
SCHEDULE = SchedulingHelper.arrange_schedule_time(DAG_ID)
4143
_DISK_SIZE_INCREMENT = 100
4244

45+
PRE_TEST_TIMEOUT = datetime.timedelta(minutes=10)
46+
POST_TEST_TIMEOUT = datetime.timedelta(minutes=10)
47+
TEST_TIMEOUT = DAGRUN_TIMEOUT - PRE_TEST_TIMEOUT - POST_TEST_TIMEOUT
48+
4349
# Keyword arguments are generated dynamically at runtime (pylint does not
4450
# know this signature).
4551
with models.DAG( # pylint: disable=unexpected-keyword-arg
@@ -108,76 +114,100 @@
108114

109115
jobset_name = jobset.generate_jobset_name(jobset_config.dag_id_prefix)
110116

111-
create_node_pool = node_pool.create.override(task_id="create_node_pool")(
112-
node_pool=cluster_info,
113-
node_pool_selector=selector,
114-
)
115-
116-
startup = jobset.create_jobset_startup_tasks(
117-
node_pool=cluster_info,
118-
jobset_config=jobset_config,
119-
jobset_name=jobset_name,
120-
node_pool_selector=selector,
121-
workload_type=Workload.JAX_TPU_BENCHMARK,
122-
)
123-
124-
node_pool_resize_start_time = node_pool.update.override(
125-
task_id="node_pool_resize"
126-
)(
127-
node_pool=cluster_info,
128-
spec=node_pool.NodePoolUpdateSpec.DiskSize(
129-
delta=_DISK_SIZE_INCREMENT
130-
),
131-
)
132-
133-
wait_for_recovery = jobset.wait_for_jobset_recovered.override(
134-
task_id="wait_for_recovery"
135-
)(
136-
node_pool=cluster_info,
137-
jobset_config=jobset_config,
138-
jobset_name=jobset_name,
139-
)
140-
141-
verify_duration = jobset.verify_recovery_duration.override(
142-
task_id="verify_recovery_duration"
143-
)(
144-
start_time=node_pool_resize_start_time,
145-
end_time=wait_for_recovery,
146-
)
147-
148-
wait_for_metric_upload = jobset.wait_for_jobset_ttr_to_be_found.override(
149-
task_id="wait_for_jobset_ttr_to_be_found",
150-
)(
151-
node_pool=cluster_info,
152-
jobset_name=jobset_name,
153-
start_time=node_pool_resize_start_time,
154-
)
155-
156-
cleanup_workload = jobset.end_workload.override(
157-
task_id="cleanup_workload", trigger_rule=TriggerRule.ALL_DONE
158-
)(
159-
node_pool=cluster_info,
160-
jobset_config=jobset_config,
161-
jobset_name=jobset_name,
162-
).as_teardown(
163-
setups=startup.jobset_start_time
164-
)
165-
166-
cleanup_node_pool = node_pool.delete.override(
167-
task_id="cleanup_node_pool", trigger_rule=TriggerRule.ALL_DONE
168-
)(node_pool=cluster_info).as_teardown(
169-
setups=create_node_pool,
170-
)
117+
with TaskGroupWithTimeout(
118+
group_id="pre_test",
119+
timeout=PRE_TEST_TIMEOUT,
120+
) as pre_test:
121+
create_node_pool = node_pool.create.override(
122+
task_id="create_node_pool"
123+
)(
124+
node_pool=cluster_info,
125+
node_pool_selector=selector,
126+
)
127+
128+
startup = jobset.create_jobset_startup_tasks(
129+
node_pool=cluster_info,
130+
jobset_config=jobset_config,
131+
jobset_name=jobset_name,
132+
node_pool_selector=selector,
133+
workload_type=Workload.JAX_TPU_BENCHMARK,
134+
)
135+
136+
with TaskGroupWithTimeout(
137+
group_id="test",
138+
timeout=TEST_TIMEOUT,
139+
) as test:
140+
node_pool_resize_start_time = node_pool.update.override(
141+
task_id="node_pool_resize"
142+
)(
143+
node_pool=cluster_info,
144+
spec=node_pool.NodePoolUpdateSpec.DiskSize(
145+
delta=_DISK_SIZE_INCREMENT
146+
),
147+
)
148+
149+
wait_for_recovery = jobset.wait_for_jobset_recovered.override(
150+
task_id="wait_for_recovery"
151+
)(
152+
node_pool=cluster_info,
153+
jobset_config=jobset_config,
154+
jobset_name=jobset_name,
155+
)
156+
157+
verify_duration = jobset.verify_recovery_duration.override(
158+
task_id="verify_recovery_duration"
159+
)(
160+
start_time=node_pool_resize_start_time,
161+
end_time=wait_for_recovery,
162+
)
163+
164+
wait_for_metric_upload = (
165+
jobset.wait_for_jobset_ttr_to_be_found.override(
166+
task_id="wait_for_jobset_ttr_to_be_found",
167+
)(
168+
node_pool=cluster_info,
169+
jobset_name=jobset_name,
170+
start_time=node_pool_resize_start_time,
171+
)
172+
)
173+
174+
chain(
175+
node_pool_resize_start_time,
176+
wait_for_recovery,
177+
verify_duration,
178+
wait_for_metric_upload,
179+
)
180+
181+
with TaskGroupWithTimeout(
182+
group_id="post_test",
183+
timeout=POST_TEST_TIMEOUT,
184+
is_teardown=True,
185+
) as post_test:
186+
cleanup_workload = jobset.end_workload.override(
187+
task_id="cleanup_workload", trigger_rule=TriggerRule.ALL_DONE
188+
)(
189+
node_pool=cluster_info,
190+
jobset_config=jobset_config,
191+
jobset_name=jobset_name,
192+
).as_teardown(
193+
setups=startup.jobset_start_time
194+
)
195+
196+
cleanup_node_pool = node_pool.delete.override(
197+
task_id="cleanup_node_pool", trigger_rule=TriggerRule.ALL_DONE
198+
)(node_pool=cluster_info).as_teardown(
199+
setups=create_node_pool,
200+
)
201+
202+
chain(
203+
cleanup_workload,
204+
cleanup_node_pool,
205+
)
171206

172207
chain(
173208
selector,
174209
jobset_name,
175-
create_node_pool,
176-
*startup.tasks,
177-
node_pool_resize_start_time,
178-
wait_for_recovery,
179-
verify_duration,
180-
wait_for_metric_upload,
181-
cleanup_workload,
182-
cleanup_node_pool,
210+
pre_test,
211+
test,
212+
post_test,
183213
)

0 commit comments

Comments
 (0)