Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand All @@ -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):
Expand Down Expand Up @@ -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",
Expand All @@ -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):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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(
Expand Down