Skip to content

Commit 6a3ebf1

Browse files
fix: Separate config construction from Airflow task execution (GoogleCloudPlatform#1278)
Move JobSets and node-pool configuration out from task, so that during CI/CD process, GitHub action machine can access GCP bucket data. Co-authored-by: Kim <kim.pai@cienet.com>
1 parent 088bbd7 commit 6a3ebf1

18 files changed

Lines changed: 58 additions & 167 deletions

.github/workflows/dag-check.yml

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -43,6 +43,12 @@ jobs:
4343
- name: Install Python dependencies
4444
run: pip install -r .github/requirements.txt
4545

46+
- name: Authenticate with gcloud
47+
uses: google-github-actions/auth@v2
48+
with:
49+
credentials_json: ${{ secrets.GCS_SECRET }}
50+
create_credentials_file: true
51+
4652
- name: Run DAGs
4753
run: |
4854
scripts/dag-check.sh

dags/tpu_observability/jobset_ttr_drain_restart.py

Lines changed: 12 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -21,7 +21,6 @@
2121
from airflow.utils.trigger_rule import TriggerRule
2222
from airflow.utils.task_group import TaskGroup
2323

24-
2524
from airflow.decorators import task
2625

2726
from dags import composer_env
@@ -76,7 +75,7 @@ def check_nodes_number(
7675
# Keyword arguments are generated dynamically at runtime (pylint does not
7776
# know this signature).
7877
with models.DAG( # pylint: disable=unexpected-keyword-arg
79-
dag_id="jobset_ttr_drain_restart",
78+
dag_id=DAG_ID,
8079
start_date=datetime.datetime(2025, 8, 10),
8180
schedule=SCHEDULE if composer_env.is_prod_env() else None,
8281
dagrun_timeout=DAGRUN_TIMEOUT,
@@ -125,38 +124,33 @@ def check_nodes_number(
125124
with TaskGroup( # pylint: disable=unexpected-keyword-arg
126125
group_id=f"v{config.tpu_version.value}"
127126
):
127+
selector = jobset.generate_node_pool_selector(DAG_ID)
128+
128129
jobset_config = jobset.build_jobset_from_gcs_yaml(
129130
gcs_path=GCS_JOBSET_CONFIG_PATH,
130-
dag_name="jobset_ttr_drain_restart",
131+
dag_name=DAG_ID,
132+
node_pool_selector=selector,
131133
)
132134

133-
cluster_info = node_pool.build_node_pool_info_from_gcs_yaml.override(
134-
task_id="build_node_pool_info_from_gcs_yaml"
135-
)(
135+
cluster_info = node_pool.build_node_pool_info_from_gcs_yaml(
136136
gcs_path=GCS_CONFIG_PATH,
137-
dag_name="jobset_ttr_drain_restart",
137+
dag_name=DAG_ID,
138138
is_prod=composer_env.is_prod_env(),
139139
machine_type=config.machine_version.value,
140140
tpu_topology=config.tpu_topology,
141+
node_pool_selector=selector,
141142
)
142143

143144
create_node_pool = node_pool.create.override(task_id="create_node_pool")(
144145
node_pool=cluster_info,
145146
)
146147

147-
start_workload = jobset.run_workload.override(task_id="start_workload")(
148+
startup = jobset.create_jobset_startup_tasks(
148149
node_pool=cluster_info,
149150
jobset_config=jobset_config,
150151
workload_type=Workload.JAX_TPU_BENCHMARK,
151152
)
152153

153-
ensure_all_pods_running = jobset.wait_for_all_pods_running.override(
154-
task_id="ensure_all_pods_running"
155-
)(
156-
node_pool=cluster_info,
157-
jobset_config=jobset_config,
158-
)
159-
160154
select_node = node_pool.draw_random_node.override(task_id="select_node")(
161155
node_pool=cluster_info
162156
)
@@ -193,7 +187,7 @@ def check_nodes_number(
193187
node_pool=cluster_info,
194188
jobset_config=jobset_config,
195189
).as_teardown(
196-
setups=start_workload
190+
setups=startup.jobset_start_time
197191
)
198192

199193
cleanup_node_pool = node_pool.delete.override(
@@ -203,11 +197,9 @@ def check_nodes_number(
203197
)
204198

205199
chain(
206-
jobset_config,
207-
cluster_info,
200+
selector,
208201
create_node_pool,
209-
start_workload,
210-
ensure_all_pods_running,
202+
*startup.tasks,
211203
select_node,
212204
drained_node,
213205
check_nodes_number,

dags/tpu_observability/jobset_ttr_kill_process.py

Lines changed: 1 addition & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -135,9 +135,7 @@ def kill_tpu_pod_workload(info: node_pool.Info, pod_name: str) -> None:
135135
node_pool_selector=selector,
136136
)
137137

138-
cluster_info = node_pool.build_node_pool_info_from_gcs_yaml.override(
139-
task_id="build_node_pool_info_from_gcs_yaml"
140-
)(
138+
cluster_info = node_pool.build_node_pool_info_from_gcs_yaml(
141139
gcs_path=GCS_CONFIG_PATH,
142140
dag_name=DAG_ID,
143141
is_prod=composer_env.is_prod_env(),
@@ -186,8 +184,6 @@ def kill_tpu_pod_workload(info: node_pool.Info, pod_name: str) -> None:
186184

187185
chain(
188186
selector,
189-
jobset_config,
190-
cluster_info,
191187
create_node_pool,
192188
*startup.tasks,
193189
kill_tasks,

dags/tpu_observability/jobset_ttr_node_pool_resize.py

Lines changed: 1 addition & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -99,9 +99,7 @@
9999
node_pool_selector=selector,
100100
)
101101

102-
cluster_info = node_pool.build_node_pool_info_from_gcs_yaml.override(
103-
task_id="build_node_pool_info_from_gcs_yaml"
104-
)(
102+
cluster_info = node_pool.build_node_pool_info_from_gcs_yaml(
105103
gcs_path=GCS_CONFIG_PATH,
106104
dag_name=DAG_ID,
107105
is_prod=composer_env.is_prod_env(),
@@ -151,8 +149,6 @@
151149

152150
chain(
153151
selector,
154-
jobset_config,
155-
cluster_info,
156152
create_node_pool,
157153
*startup.tasks,
158154
node_pool_resize,

dags/tpu_observability/jobset_ttr_node_reboot.py

Lines changed: 1 addition & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -104,9 +104,7 @@
104104
privileged=True,
105105
)
106106

107-
cluster_info = node_pool.build_node_pool_info_from_gcs_yaml.override(
108-
task_id="build_node_pool_info_from_gcs_yaml"
109-
)(
107+
cluster_info = node_pool.build_node_pool_info_from_gcs_yaml(
110108
gcs_path=GCS_CONFIG_PATH,
111109
dag_name=DAG_ID,
112110
is_prod=composer_env.is_prod_env(),
@@ -158,8 +156,6 @@
158156

159157
chain(
160158
selector,
161-
jobset_config,
162-
cluster_info,
163159
create_node_pool,
164160
*startup.tasks,
165161
target_pod,

dags/tpu_observability/jobset_ttr_pod_delete.py

Lines changed: 1 addition & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -95,9 +95,7 @@
9595
node_pool_selector=selector,
9696
)
9797

98-
cluster_info = node_pool.build_node_pool_info_from_gcs_yaml.override(
99-
task_id="build_node_pool_info_from_gcs_yaml"
100-
)(
98+
cluster_info = node_pool.build_node_pool_info_from_gcs_yaml(
10199
gcs_path=GCS_CONFIG_PATH,
102100
dag_name=DAG_ID,
103101
is_prod=composer_env.is_prod_env(),
@@ -144,8 +142,6 @@
144142

145143
chain(
146144
selector,
147-
jobset_config,
148-
cluster_info,
149145
create_node_pool,
150146
*startup.tasks,
151147
delete_random_pod,

dags/tpu_observability/jobset_ttr_rollback.py

Lines changed: 1 addition & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -96,9 +96,7 @@
9696
node_pool_selector=selector,
9797
)
9898

99-
cluster_info = node_pool.build_node_pool_info_from_gcs_yaml.override(
100-
task_id="build_node_pool_info_from_gcs_yaml"
101-
)(
99+
cluster_info = node_pool.build_node_pool_info_from_gcs_yaml(
102100
gcs_path=GCS_CONFIG_PATH,
103101
dag_name=DAG_ID,
104102
is_prod=composer_env.is_prod_env(),
@@ -145,8 +143,6 @@
145143

146144
chain(
147145
selector,
148-
jobset_config,
149-
cluster_info,
150146
create_node_pool,
151147
*startup.tasks,
152148
rollback_node_pool,

dags/tpu_observability/jobset_uptime_validation.py

Lines changed: 1 addition & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -107,9 +107,7 @@ def get_current_time() -> TimeUtil:
107107
node_pool_selector=selector,
108108
)
109109

110-
cluster_info = node_pool.build_node_pool_info_from_gcs_yaml.override(
111-
task_id="build_node_pool_info_from_gcs_yaml"
112-
)(
110+
cluster_info = node_pool.build_node_pool_info_from_gcs_yaml(
113111
gcs_path=GCS_CONFIG_PATH,
114112
dag_name=DAG_ID,
115113
is_prod=composer_env.is_prod_env(),
@@ -167,8 +165,6 @@ def get_current_time() -> TimeUtil:
167165

168166
chain(
169167
selector,
170-
jobset_config,
171-
cluster_info,
172168
create_node_pool,
173169
*startup.tasks,
174170
wait_for_jobset_uptime_data,

dags/tpu_observability/multi_host_nodepool_rollback_dag.py

Lines changed: 1 addition & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -87,9 +87,7 @@
8787
with TaskGroup( # pylint: disable=unexpected-keyword-arg
8888
group_id=f"v{config.tpu_version.value}"
8989
):
90-
node_pool_info = node_pool.build_node_pool_info_from_gcs_yaml.override(
91-
task_id="build_node_pool_info_from_gcs_yaml"
92-
)(
90+
node_pool_info = node_pool.build_node_pool_info_from_gcs_yaml(
9391
gcs_path=GCS_CONFIG_PATH,
9492
dag_name=DAG_ID,
9593
is_prod=composer_env.is_prod_env(),
@@ -125,7 +123,6 @@
125123
)
126124

127125
chain(
128-
node_pool_info,
129126
create_node_pool,
130127
wait_node_pool_available,
131128
rollback_node_pool,

dags/tpu_observability/node_pool_status.py

Lines changed: 6 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -15,9 +15,9 @@
1515
"""A DAG to validate the status of a GKE node pool through its lifecycle."""
1616

1717
import datetime
18+
import copy
1819

1920
from airflow import models
20-
from airflow.decorators import task
2121
from airflow.models.baseoperator import chain
2222
from airflow.utils.task_group import TaskGroup
2323
from airflow.utils.trigger_rule import TriggerRule
@@ -69,39 +69,23 @@
6969
for machine in MachineConfigMap:
7070
config = machine.value
7171

72-
@task
73-
def generate_problematic_node_pool_name(
74-
node_pool_info: node_pool.Info,
75-
) -> str:
76-
"""Generates a problematic node pool name."""
77-
return f"{node_pool_info.node_pool_name}-x"
78-
79-
@task
80-
def generate_problematic_node_location(
81-
node_pool_info: node_pool.Info,
82-
) -> str:
83-
"""Generates a problematic node location."""
84-
return f"{node_pool_info.location}-c"
85-
8672
# Keyword arguments are generated dynamically at runtime (pylint does not
8773
# know this signature).
8874
with TaskGroup( # pylint: disable=unexpected-keyword-arg
8975
group_id=f"v{config.tpu_version.value}"
9076
):
91-
node_pool_info = node_pool.build_node_pool_info_from_gcs_yaml.override(
92-
task_id="build_node_pool_info_from_gcs_yaml"
93-
)(
77+
node_pool_info = node_pool.build_node_pool_info_from_gcs_yaml(
9478
gcs_path=GCS_CONFIG_PATH,
9579
dag_name=DAG_ID,
9680
is_prod=composer_env.is_prod_env(),
9781
machine_type=config.machine_version.value,
9882
tpu_topology=config.tpu_topology,
9983
)
10084

101-
problematic_node_pool_info = node_pool.copy_node_pool_info_with_override(
102-
info=node_pool_info,
103-
node_pool_name=generate_problematic_node_pool_name(node_pool_info),
104-
node_locations=generate_problematic_node_location(node_pool_info),
85+
problematic_node_pool_info = copy.deepcopy(node_pool_info)
86+
problematic_node_pool_info.location = f"{node_pool_info.location}-c"
87+
problematic_node_pool_info.node_pool_name = (
88+
f"{node_pool_info.node_pool_name}-x"
10589
)
10690

10791
task_id = "create_node_pool"
@@ -191,8 +175,6 @@ def generate_problematic_node_location(
191175
)
192176

193177
chain(
194-
node_pool_info,
195-
problematic_node_pool_info,
196178
create_node_pool,
197179
wait_for_provisioning,
198180
wait_for_running,

0 commit comments

Comments
 (0)