|
21 | 21 | from airflow.models.baseoperator import chain |
22 | 22 | from airflow.utils.task_group import TaskGroup |
23 | 23 | 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 |
24 | 29 |
|
25 | 30 | from dags import composer_env |
26 | 31 | from dags.common.scheduling_helper.scheduling_helper import ( |
|
42 | 47 | DAGRUN_TIMEOUT = get_dag_timeout(DAG_ID) |
43 | 48 | SCHEDULE = SchedulingHelper.arrange_schedule_time(DAG_ID) |
44 | 49 |
|
| 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 | + |
45 | 54 |
|
46 | 55 | @task |
47 | 56 | def check_nodes_number( |
@@ -140,75 +149,102 @@ def check_nodes_number( |
140 | 149 | selector = jobset.generate_node_pool_selector(DAG_ID) |
141 | 150 | jobset_name = jobset.generate_jobset_name(jobset_config.dag_id_prefix) |
142 | 151 |
|
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 | + ) |
201 | 243 |
|
202 | 244 | chain( |
203 | 245 | selector, |
204 | 246 | 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, |
214 | 250 | ) |
0 commit comments