diff --git a/providers/google/src/airflow/providers/google/cloud/operators/cloud_batch.py b/providers/google/src/airflow/providers/google/cloud/operators/cloud_batch.py index 914559070b7b2..83992a09eae7f 100644 --- a/providers/google/src/airflow/providers/google/cloud/operators/cloud_batch.py +++ b/providers/google/src/airflow/providers/google/cloud/operators/cloud_batch.py @@ -78,10 +78,6 @@ def __init__( self.region = region self.job_name = job_name self.job = job - # Normalize Job protobuf to dict so Airflow's template renderer can descend - # into nested fields (e.g. runnable.container.commands). See #37217. - if isinstance(job, Job): - self.job = Job.to_dict(job) self.polling_period_seconds = polling_period_seconds self.timeout_seconds = timeout_seconds self.gcp_conn_id = gcp_conn_id @@ -89,6 +85,13 @@ def __init__( self.deferrable = deferrable self.polling_period_seconds = polling_period_seconds + def prepare_template(self) -> None: + # Normalize Job protobuf to dict so Airflow's template renderer can descend + # into nested fields (e.g. runnable.container.commands) before rendering runs. + # See #37217. + if isinstance(self.job, Job): + self.job = Job.to_dict(self.job) + def execute(self, context: Context): hook: CloudBatchHook = CloudBatchHook(self.gcp_conn_id, self.impersonation_chain) job = hook.submit_batch_job( diff --git a/providers/google/tests/unit/google/cloud/operators/test_cloud_batch.py b/providers/google/tests/unit/google/cloud/operators/test_cloud_batch.py index 1b688bb65dfd0..44436b3f99080 100644 --- a/providers/google/tests/unit/google/cloud/operators/test_cloud_batch.py +++ b/providers/google/tests/unit/google/cloud/operators/test_cloud_batch.py @@ -48,6 +48,7 @@ def test_execute(self, mock): operator = CloudBatchSubmitJobOperator( task_id=TASK_ID, project_id=PROJECT_ID, region=REGION, job_name=JOB_NAME, job=JOB ) + operator.prepare_template() completed_job = operator.execute(context=mock.MagicMock()) @@ -119,6 +120,18 @@ class TestCloudBatchSubmitJobOperatorTemplating: def test_template_fields_includes_job(self): assert "job" in CloudBatchSubmitJobOperator.template_fields + def test_protobuf_job_is_normalized_by_prepare_template_not_init(self): + job = batch_v1.Job.from_json(json.dumps(_job_dict_with_template())) + operator = CloudBatchSubmitJobOperator( + task_id=TASK_ID, project_id=PROJECT_ID, region=REGION, job_name=JOB_NAME, job=job + ) + + assert isinstance(operator.job, batch_v1.Job) + + operator.prepare_template() + + assert isinstance(operator.job, dict) + @pytest.mark.db_test @pytest.mark.parametrize( "job_input_factory", diff --git a/scripts/ci/prek/validate_operators_init_exemptions.txt b/scripts/ci/prek/validate_operators_init_exemptions.txt index e6869670a5bea..210e077095267 100644 --- a/scripts/ci/prek/validate_operators_init_exemptions.txt +++ b/scripts/ci/prek/validate_operators_init_exemptions.txt @@ -9,7 +9,6 @@ providers/amazon/src/airflow/providers/amazon/aws/operators/neptune.py::NeptuneS providers/amazon/src/airflow/providers/amazon/aws/operators/neptune.py::NeptuneStopDbClusterOperator providers/amazon/src/airflow/providers/amazon/aws/transfers/gcs_to_s3.py::GCSToS3Operator providers/cncf/kubernetes/src/airflow/providers/cncf/kubernetes/operators/pod.py::KubernetesPodOperator -providers/google/src/airflow/providers/google/cloud/operators/cloud_batch.py::CloudBatchSubmitJobOperator providers/google/src/airflow/providers/google/cloud/operators/cloud_build.py::CloudBuildCreateBuildOperator providers/google/src/airflow/providers/google/cloud/operators/cloud_storage_transfer_service.py::CloudDataTransferServiceCreateJobOperator providers/google/src/airflow/providers/google/cloud/operators/dataproc.py::DataprocCreateClusterOperator