From 7e9789bf0965c0c3ec1e6f67482582728471f168 Mon Sep 17 00:00:00 2001 From: PoAn Yang Date: Sat, 12 Sep 2026 19:40:04 +0900 Subject: [PATCH] Fail Gemini batch job operators when the job does not succeed Signed-off-by: PoAn Yang --- .../google/cloud/operators/gen_ai.py | 6 +- .../google/cloud/operators/test_gen_ai.py | 58 +++++++++++++++++++ 2 files changed, 62 insertions(+), 2 deletions(-) diff --git a/providers/google/src/airflow/providers/google/cloud/operators/gen_ai.py b/providers/google/src/airflow/providers/google/cloud/operators/gen_ai.py index 65d8a9433742c..2b06bd87f7a22 100644 --- a/providers/google/src/airflow/providers/google/cloud/operators/gen_ai.py +++ b/providers/google/src/airflow/providers/google/cloud/operators/gen_ai.py @@ -487,7 +487,6 @@ def _wait_until_complete(self, job, polling_interval: int = 30): BatchJobStatus.EXPIRED.value, BatchJobStatus.CANCELLED.value, ]: - self.log.error("Job execution was not completed!") break self.log.info( "Waiting for job execution, polling interval: %s seconds, current state: %s", @@ -497,6 +496,8 @@ def _wait_until_complete(self, job, polling_interval: int = 30): time.sleep(polling_interval) except Exception: raise AirflowException("Something went wrong during waiting of the batch job.") + if job.state.name != BatchJobStatus.SUCCEEDED.value: + raise RuntimeError(f"Job {job.name} execution was not completed! state: {job.state.name}") return job def _validate_results_folder(self): @@ -937,7 +938,6 @@ def _wait_until_complete(self, job, polling_interval: int = 30): BatchJobStatus.EXPIRED.value, BatchJobStatus.CANCELLED.value, ]: - self.log.error("Job execution was not completed!") break self.log.info( "Waiting for job execution, polling interval: %s seconds, current state: %s", @@ -947,6 +947,8 @@ def _wait_until_complete(self, job, polling_interval: int = 30): time.sleep(polling_interval) except Exception as e: raise AirflowException("Something went wrong during waiting of the batch job: %s", e) + if job.state.name != BatchJobStatus.SUCCEEDED.value: + raise RuntimeError(f"Job {job.name} execution was not completed! state: {job.state.name}") return job def _validate_results_folder(self): diff --git a/providers/google/tests/unit/google/cloud/operators/test_gen_ai.py b/providers/google/tests/unit/google/cloud/operators/test_gen_ai.py index 1071f2de82920..770abbd66f4ed 100644 --- a/providers/google/tests/unit/google/cloud/operators/test_gen_ai.py +++ b/providers/google/tests/unit/google/cloud/operators/test_gen_ai.py @@ -21,10 +21,12 @@ import pytest from google.genai.errors import ClientError from google.genai.types import ( + BatchJob, Content, CreateCachedContentConfig, GenerateContentConfig, GoogleSearch, + JobState, Part, Tool, TuningDataset, @@ -500,6 +502,34 @@ def test__wait_until_complete_exception_raises_airflow_exception(self, mock_hook with pytest.raises(AirflowException): op._wait_until_complete(job=mock.MagicMock()) + @pytest.mark.parametrize( + "job_state", + [JobState.JOB_STATE_FAILED, JobState.JOB_STATE_EXPIRED, JobState.JOB_STATE_CANCELLED], + ) + @mock.patch(GEN_AI_PATH.format("GenAIGeminiAPIHook"), autospec=True) + def test_execute_wait_until_complete_unsuccessful_job_raises_runtime_error(self, mock_hook, job_state): + mock_hook.return_value.get_batch_job.return_value = BatchJob( + name=TEST_BATCH_JOB_NAME, state=job_state + ) + op = GenAIGeminiCreateBatchJobOperator( + task_id=TASK_ID, + project_id=GCP_PROJECT, + location=GCP_LOCATION, + model=TEST_GEMINI_MODEL, + gcp_conn_id=GCP_CONN_ID, + impersonation_chain=IMPERSONATION_CHAIN, + input_source=TEST_BATCH_JOB_INLINED_REQUESTS, + gemini_api_key=TEST_GEMINI_API_KEY, + wait_until_complete=True, + deferrable=False, + ) + + with pytest.raises( + RuntimeError, + match=f"Job {TEST_BATCH_JOB_NAME} execution was not completed! state: {job_state.name}", + ): + op.execute(context={"ti": mock.Mock(spec_set=["xcom_push"])}) + @mock.patch(GEN_AI_PATH.format("GenAIGeminiAPIHook")) def test_execute_exception_error_raises_airflow_exception(self, mock_hook): op = GenAIGeminiCreateBatchJobOperator( @@ -900,6 +930,34 @@ def test__wait_until_complete_exception_raises_airflow_exception(self, mock_hook with pytest.raises(AirflowException): op._wait_until_complete(job=mock.MagicMock()) + @pytest.mark.parametrize( + "job_state", + [JobState.JOB_STATE_FAILED, JobState.JOB_STATE_EXPIRED, JobState.JOB_STATE_CANCELLED], + ) + @mock.patch(GEN_AI_PATH.format("GenAIGeminiAPIHook"), autospec=True) + def test_execute_wait_until_complete_unsuccessful_job_raises_runtime_error(self, mock_hook, job_state): + mock_hook.return_value.get_batch_job.return_value = BatchJob( + name=TEST_BATCH_JOB_NAME, state=job_state + ) + op = GenAIGeminiCreateEmbeddingsBatchJobOperator( + task_id=TASK_ID, + project_id=GCP_PROJECT, + location=GCP_LOCATION, + input_source=TEST_EMBEDDINGS_JOB_INLINED_REQUESTS, + model=EMBEDDING_MODEL, + gemini_api_key=TEST_GEMINI_API_KEY, + gcp_conn_id=GCP_CONN_ID, + impersonation_chain=IMPERSONATION_CHAIN, + wait_until_complete=True, + deferrable=False, + ) + + with pytest.raises( + RuntimeError, + match=f"Job {TEST_BATCH_JOB_NAME} execution was not completed! state: {job_state.name}", + ): + op.execute(context={"ti": mock.Mock(spec_set=["xcom_push"])}) + @mock.patch(GEN_AI_PATH.format("GenAIGeminiAPIHook")) def test_execute_exception_error_raises_airflow_exception(self, mock_hook): op = GenAIGeminiCreateEmbeddingsBatchJobOperator(