Skip to content

Commit 08fa8b2

Browse files
authored
feat: Apply TaskGroupWithTimeout for TPU Observability DAGs - part 1
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 505afa9 commit 08fa8b2

17 files changed

Lines changed: 887 additions & 594 deletions

.github/workflows/dag-check.yml

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -3,7 +3,6 @@ name: DAG Check
33

44
on:
55
pull_request:
6-
branches: [master]
76
types: [opened, synchronize, edited]
87

98
push:

.github/workflows/pyink-check.yml

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -2,11 +2,12 @@ name: Formatter
22

33
on:
44
pull_request:
5-
branches: [master]
65
types: [opened, synchronize, edited]
76
push:
87
branches: [master]
98

9+
workflow_dispatch: {}
10+
1011
jobs:
1112
format_check:
1213
runs-on: ubuntu-latest

.github/workflows/pylint-check.yml

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -2,12 +2,13 @@ name: Linter
22

33
on:
44
pull_request:
5-
branches: [master]
65
types: [opened, synchronize, edited]
76

87
push:
98
branches: [master]
109

10+
workflow_dispatch: {}
11+
1112
jobs:
1213
linting_check:
1314
runs-on: ubuntu-latest

.github/workflows/require-checklist.yml

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -2,10 +2,13 @@ name: Require Checklist
22
on:
33
pull_request:
44
types: [opened, edited, synchronize]
5+
6+
workflow_dispatch: {}
7+
58
jobs:
69
check_pr_body:
710
runs-on: ubuntu-latest
811
steps:
912
- uses: mheap/require-checklist-action@v2
1013
with:
11-
requireChecklist: false # If this is true and there are no checklists detected, the action will fail
14+
requireChecklist: false # If this is true and there are no checklists detected, the action will fail

.github/workflows/unit-test.yml

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -3,7 +3,6 @@ name: Unit Test
33

44
on:
55
pull_request:
6-
branches: [master]
76
types: [opened, synchronize, edited]
87

98
push:

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
)

0 commit comments

Comments
 (0)