|
74 | 74 | _TEST_PIPELINE_JOB_DISPLAY_NAME_2 = "sample-pipeline-job-display-name-2" |
75 | 75 | _TEST_PIPELINE_JOB_ID = "sample-test-pipeline-202111111" |
76 | 76 | _TEST_PIPELINE_JOB_ID_2 = "sample-test-pipeline-202111112" |
| 77 | +_TEST_PIPELINE_RERUN_JOB_ID = "sample-test-pipeline-rerun" |
77 | 78 | _TEST_GCS_BUCKET_NAME = "my-bucket" |
78 | 79 | _TEST_GCS_OUTPUT_DIRECTORY = f"gs://{_TEST_GCS_BUCKET_NAME}/output_artifacts/" |
79 | 80 | _TEST_CREDENTIALS = auth_credentials.AnonymousCredentials() |
@@ -339,6 +340,41 @@ def mock_pipeline_v1beta1_service_get(): |
339 | 340 | yield mock_get_pipeline_job |
340 | 341 |
|
341 | 342 |
|
| 343 | +@pytest.fixture |
| 344 | +def mock_pipeline_v1beta1_service_get_with_task_details(): |
| 345 | + pipeline_job = make_v1beta1_pipeline_job( |
| 346 | + _TEST_PIPELINE_JOB_NAME, |
| 347 | + gca_pipeline_state.PipelineState.PIPELINE_STATE_FAILED, |
| 348 | + ) |
| 349 | + pipeline_job.job_detail.task_details = [ |
| 350 | + v1beta1_pipeline_job.PipelineTaskDetail( |
| 351 | + task_id=1, |
| 352 | + task_name="task-1", |
| 353 | + state=aiplatform_v1beta1.PipelineTaskDetail.State.SUCCEEDED, |
| 354 | + task_unique_name="task-1", |
| 355 | + ), |
| 356 | + v1beta1_pipeline_job.PipelineTaskDetail( |
| 357 | + task_id=2, |
| 358 | + task_name="task-2", |
| 359 | + state=aiplatform_v1beta1.PipelineTaskDetail.State.FAILED, |
| 360 | + task_unique_name="task-2", |
| 361 | + ), |
| 362 | + v1beta1_pipeline_job.PipelineTaskDetail( |
| 363 | + task_id=3, |
| 364 | + task_name="task-3", |
| 365 | + state=aiplatform_v1beta1.PipelineTaskDetail.State.SUCCEEDED, |
| 366 | + task_unique_name="task-3", |
| 367 | + ), |
| 368 | + ] |
| 369 | + |
| 370 | + with mock.patch.object( |
| 371 | + v1beta1_pipeline_service.PipelineServiceClient, "get_pipeline_job" |
| 372 | + ) as mock_get_pipeline_job: |
| 373 | + mock_get_pipeline_job.side_effect = [pipeline_job] |
| 374 | + |
| 375 | + yield mock_get_pipeline_job |
| 376 | + |
| 377 | + |
342 | 378 | @pytest.fixture |
343 | 379 | def mock_pipeline_v1_service_batch_cancel(): |
344 | 380 | with patch.object( |
@@ -2417,6 +2453,63 @@ def test_rerun_v1beta1_pipeline_job_returns_response( |
2417 | 2453 | assert mock_pipeline_v1beta1_service_get.call_count == 1 |
2418 | 2454 | assert mock_pipeline_v1beta1_service_create.call_count == 2 |
2419 | 2455 |
|
| 2456 | + @pytest.mark.usefixtures( |
| 2457 | + "mock_pipeline_v1beta1_service_create", |
| 2458 | + "mock_pipeline_v1beta1_service_get_with_task_details", |
| 2459 | + ) |
| 2460 | + @pytest.mark.parametrize( |
| 2461 | + "job_spec", |
| 2462 | + [_TEST_PIPELINE_SPEC_JSON, _TEST_PIPELINE_SPEC_YAML, _TEST_PIPELINE_JOB], |
| 2463 | + ) |
| 2464 | + def test_rerun_v1beta1_pipeline_job_returns_response_with_task_details( |
| 2465 | + self, |
| 2466 | + mock_load_yaml_and_json, |
| 2467 | + job_spec, |
| 2468 | + mock_pipeline_v1beta1_service_create, |
| 2469 | + mock_pipeline_v1beta1_service_get_with_task_details, |
| 2470 | + ): |
| 2471 | + aiplatform.init( |
| 2472 | + project=_TEST_PROJECT, |
| 2473 | + staging_bucket=_TEST_GCS_BUCKET_NAME, |
| 2474 | + credentials=_TEST_CREDENTIALS, |
| 2475 | + ) |
| 2476 | + |
| 2477 | + job = preview_pipeline_jobs._PipelineJob( |
| 2478 | + display_name=_TEST_PIPELINE_JOB_DISPLAY_NAME, |
| 2479 | + template_path=_TEST_TEMPLATE_PATH, |
| 2480 | + job_id=_TEST_PIPELINE_JOB_ID, |
| 2481 | + ) |
| 2482 | + |
| 2483 | + job.submit() |
| 2484 | + |
| 2485 | + job.rerun( |
| 2486 | + job_id=_TEST_PIPELINE_RERUN_JOB_ID, |
| 2487 | + original_pipelinejob_name=_TEST_PIPELINE_JOB_NAME, |
| 2488 | + pipeline_task_rerun_configs=[ |
| 2489 | + aiplatform_v1beta1.PipelineTaskRerunConfig( |
| 2490 | + task_id=1, |
| 2491 | + skip_task=False, |
| 2492 | + ) |
| 2493 | + ], |
| 2494 | + parameter_values={"param-1": "value-1"}, |
| 2495 | + ) |
| 2496 | + |
| 2497 | + assert mock_pipeline_v1beta1_service_get_with_task_details.call_count == 1 |
| 2498 | + assert mock_pipeline_v1beta1_service_create.call_count == 2 |
| 2499 | + |
| 2500 | + rerun_config = mock_pipeline_v1beta1_service_create.call_args[1][ |
| 2501 | + "request" |
| 2502 | + ].pipeline_job.pipeline_task_rerun_configs |
| 2503 | + assert rerun_config[0].task_name == "task-1" |
| 2504 | + assert not rerun_config[0].skip_task |
| 2505 | + assert rerun_config[0].task_id == 1 |
| 2506 | + assert rerun_config[1].task_name == "task-2" |
| 2507 | + assert not rerun_config[1].skip_task |
| 2508 | + assert rerun_config[1].task_id == 2 |
| 2509 | + assert rerun_config[2].task_name == "task-3" |
| 2510 | + assert rerun_config[2].skip_task |
| 2511 | + assert rerun_config[2].task_id == 3 |
| 2512 | + |
2420 | 2513 | @pytest.mark.parametrize( |
2421 | 2514 | "job_spec", |
2422 | 2515 | [_TEST_PIPELINE_SPEC_JSON, _TEST_PIPELINE_SPEC_YAML, _TEST_PIPELINE_JOB], |
|
0 commit comments