Skip to content

Commit de4f4d3

Browse files
authored
feat: Apply TaskGroupWithTimeout for TPU Observability DAGs - part 1 (GoogleCloudPlatform#1299)
Integrates `TaskGroupWithTimeout` directly into all TPU Observability DAGs that do not depend on Kubernetes JobSets. This provides strict, shared deadline timeout enforcement across distinct stages of GKE node pool validation. Specifically: 1. Configures stage-specific timeouts for all modified DAGs: - `PRE_TEST_TIMEOUT`: 10 minutes (for GKE provisioning and availability setup) - `POST_TEST_TIMEOUT`: 10 minutes (for GKE teardown and cleanup) - `TEST_TIMEOUT`: Evaluates dynamically as `DAGRUN_TIMEOUT - 20 mins` 2. Restructures the following 5 DAGs into three separate `TaskGroupWithTimeout` blocks: - `dags/tpu_observability/node_pool_status.py` - `dags/tpu_observability/node_pool_ttr_disk_size.py` - `dags/tpu_observability/node_pool_ttr_update_label.py` - `dags/tpu_observability/multi_host_nodepool_rollback_dag.py` - `dags/tpu_observability/update_node_pool_label.py` 3. Refactors GKE node pool operations into: - `pre_test`: Handles node pool creation and waiting for provisioning/running. - `test`: Executes the core test mutation and verification flow. - `post_test`: Manages cleanup/teardown of GKE node pools (using `is_teardown=True` to guarantee execution even if upstream stages timeout or fail). * 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 a6eceeb commit de4f4d3

11 files changed

Lines changed: 847 additions & 580 deletions

dags/tpu_observability/jobset_ttr_drain_restart.py

Lines changed: 103 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,102 @@ 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+
with TaskGroupWithTimeout(
164+
group_id="test",
165+
timeout=TEST_TIMEOUT,
166+
) as test:
167+
startup = jobset.create_jobset_startup_tasks(
168+
node_pool=cluster_info,
169+
jobset_config=jobset_config,
170+
jobset_name=jobset_name,
171+
node_pool_selector=selector,
172+
workload_type=Workload.JAX_TPU_BENCHMARK,
173+
)
174+
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+
*startup.tasks,
211+
select_node,
212+
drained_node,
213+
check_nodes_number_task,
214+
uncordon_node,
215+
wait_for_metric_upload,
216+
)
217+
218+
with TaskGroupWithTimeout(
219+
group_id="post_test",
220+
timeout=POST_TEST_TIMEOUT,
221+
is_teardown=True,
222+
) as post_test:
223+
cleanup_workload = jobset.end_workload.override(
224+
task_id="cleanup_workload", trigger_rule=TriggerRule.ALL_DONE
225+
)(
226+
node_pool=cluster_info,
227+
jobset_config=jobset_config,
228+
jobset_name=jobset_name,
229+
).as_teardown(
230+
setups=startup.jobset_start_time
231+
)
232+
233+
cleanup_node_pool = node_pool.delete.override(
234+
task_id="cleanup_node_pool", trigger_rule=TriggerRule.ALL_DONE
235+
)(node_pool=cluster_info).as_teardown(
236+
setups=create_node_pool,
237+
)
238+
239+
chain(
240+
cleanup_workload,
241+
cleanup_node_pool,
242+
)
201243

202244
chain(
203245
selector,
204246
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,
247+
pre_test,
248+
test,
249+
post_test,
214250
)

dags/tpu_observability/jobset_ttr_node_pool_resize.py

Lines changed: 99 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,101 @@
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+
with TaskGroupWithTimeout(
129+
group_id="test",
130+
timeout=TEST_TIMEOUT,
131+
) as test:
132+
startup = jobset.create_jobset_startup_tasks(
133+
node_pool=cluster_info,
134+
jobset_config=jobset_config,
135+
jobset_name=jobset_name,
136+
node_pool_selector=selector,
137+
workload_type=Workload.JAX_TPU_BENCHMARK,
138+
)
139+
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+
*startup.tasks,
176+
node_pool_resize_start_time,
177+
wait_for_recovery,
178+
verify_duration,
179+
wait_for_metric_upload,
180+
)
181+
182+
with TaskGroupWithTimeout(
183+
group_id="post_test",
184+
timeout=POST_TEST_TIMEOUT,
185+
is_teardown=True,
186+
) as post_test:
187+
cleanup_workload = jobset.end_workload.override(
188+
task_id="cleanup_workload", trigger_rule=TriggerRule.ALL_DONE
189+
)(
190+
node_pool=cluster_info,
191+
jobset_config=jobset_config,
192+
jobset_name=jobset_name,
193+
).as_teardown(
194+
setups=startup.jobset_start_time
195+
)
196+
197+
cleanup_node_pool = node_pool.delete.override(
198+
task_id="cleanup_node_pool", trigger_rule=TriggerRule.ALL_DONE
199+
)(node_pool=cluster_info).as_teardown(
200+
setups=create_node_pool,
201+
)
202+
203+
chain(
204+
cleanup_workload,
205+
cleanup_node_pool,
206+
)
171207

172208
chain(
173209
selector,
174210
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,
211+
pre_test,
212+
test,
213+
post_test,
183214
)

0 commit comments

Comments
 (0)