Skip to content

Commit a6b8b30

Browse files
committed
Build Dataproc submit job trigger arguments after template rendering
1 parent 6614946 commit a6b8b30

3 files changed

Lines changed: 50 additions & 7 deletions

File tree

providers/google/src/airflow/providers/google/cloud/operators/dataproc.py

Lines changed: 8 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -24,7 +24,7 @@
2424
import time
2525
import warnings
2626
from collections.abc import MutableSequence, Sequence
27-
from dataclasses import dataclass
27+
from dataclasses import dataclass, replace
2828
from datetime import datetime, timedelta
2929
from enum import Enum
3030
from functools import cached_property
@@ -66,6 +66,7 @@
6666
from airflow.triggers.base import StartTriggerArgs
6767

6868
if TYPE_CHECKING:
69+
import jinja2
6970
from google.api_core import operation
7071
from google.api_core.retry_async import AsyncRetry
7172
from google.protobuf.duration_pb2 import Duration
@@ -2004,9 +2005,13 @@ def __init__(
20042005
self.openlineage_inject_parent_job_info = openlineage_inject_parent_job_info
20052006
self.openlineage_inject_transport_info = openlineage_inject_transport_info
20062007

2008+
def render_template_fields(self, context: Context, jinja_env: jinja2.Environment | None = None) -> None:
2009+
super().render_template_fields(context, jinja_env)
2010+
# Built here rather than in __init__: every value below is a template field, so in the
2011+
# constructor they still hold the un-rendered Jinja expression.
20072012
if self.deferrable and self.start_from_trigger:
2008-
self.start_trigger_args = StartTriggerArgs(
2009-
trigger_cls="airflow.providers.google.cloud.triggers.dataproc.DataprocSubmitJobDirectTrigger",
2013+
self.start_trigger_args = replace(
2014+
self.start_trigger_args,
20102015
trigger_kwargs={
20112016
"job": self.job,
20122017
"project_id": self.project_id,
@@ -2017,9 +2022,6 @@ def __init__(
20172022
"cancel_on_kill": self.cancel_on_kill,
20182023
"request_id": self.request_id,
20192024
},
2020-
next_method="execute_complete",
2021-
next_kwargs=None,
2022-
timeout=None,
20232025
)
20242026

20252027
def execute(self, context: Context):

providers/google/tests/unit/google/cloud/operators/test_dataproc.py

Lines changed: 42 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2491,6 +2491,7 @@ def test_start_from_trigger_sets_start_trigger_args(self):
24912491
cancel_on_kill=False,
24922492
request_id=REQUEST_ID,
24932493
)
2494+
op.render_template_fields({})
24942495
assert op.start_from_trigger is True
24952496
assert (
24962497
op.start_trigger_args.trigger_cls
@@ -2508,6 +2509,46 @@ def test_start_from_trigger_sets_start_trigger_args(self):
25082509
}
25092510
assert op.start_trigger_args.next_method == "execute_complete"
25102511

2512+
def test_start_trigger_args_use_rendered_template_fields(self):
2513+
op = DataprocSubmitJobOperator(
2514+
task_id=TASK_ID,
2515+
region="{{ params.region }}",
2516+
project_id="{{ params.project }}",
2517+
job={"reference": {"job_id": "job-{{ ds }}"}},
2518+
deferrable=True,
2519+
start_from_trigger=True,
2520+
)
2521+
op.render_template_fields(
2522+
{"ds": "2026-01-01", "params": {"region": GCP_REGION, "project": GCP_PROJECT}}
2523+
)
2524+
trigger_kwargs = op.start_trigger_args.trigger_kwargs
2525+
assert trigger_kwargs["region"] == GCP_REGION
2526+
assert trigger_kwargs["project_id"] == GCP_PROJECT
2527+
assert trigger_kwargs["job"] == {"reference": {"job_id": "job-2026-01-01"}}
2528+
2529+
def test_start_trigger_args_are_not_shared_between_tasks(self):
2530+
first = DataprocSubmitJobOperator(
2531+
task_id="first",
2532+
region=GCP_REGION,
2533+
project_id=GCP_PROJECT,
2534+
job={"reference": {"job_id": "first"}},
2535+
deferrable=True,
2536+
start_from_trigger=True,
2537+
)
2538+
second = DataprocSubmitJobOperator(
2539+
task_id="second",
2540+
region=GCP_REGION,
2541+
project_id=GCP_PROJECT,
2542+
job={"reference": {"job_id": "second"}},
2543+
deferrable=True,
2544+
start_from_trigger=True,
2545+
)
2546+
first.render_template_fields({})
2547+
second.render_template_fields({})
2548+
assert first.start_trigger_args.trigger_kwargs["job"] == {"reference": {"job_id": "first"}}
2549+
assert second.start_trigger_args.trigger_kwargs["job"] == {"reference": {"job_id": "second"}}
2550+
assert DataprocSubmitJobOperator.start_trigger_args.trigger_kwargs == {}
2551+
25112552
def test_start_from_trigger_without_deferrable_does_not_set_args(self):
25122553
op = DataprocSubmitJobOperator(
25132554
task_id=TASK_ID,
@@ -2518,6 +2559,7 @@ def test_start_from_trigger_without_deferrable_does_not_set_args(self):
25182559
deferrable=False,
25192560
start_from_trigger=True,
25202561
)
2562+
op.render_template_fields({})
25212563
assert op.start_from_trigger is True
25222564
assert op.start_trigger_args.trigger_kwargs == {}
25232565

scripts/ci/prek/validate_operators_init_exemptions.txt

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -13,7 +13,6 @@ providers/google/src/airflow/providers/google/cloud/operators/cloud_batch.py::Cl
1313
providers/google/src/airflow/providers/google/cloud/operators/cloud_build.py::CloudBuildCreateBuildOperator
1414
providers/google/src/airflow/providers/google/cloud/operators/cloud_storage_transfer_service.py::CloudDataTransferServiceCreateJobOperator
1515
providers/google/src/airflow/providers/google/cloud/operators/dataproc.py::DataprocCreateClusterOperator
16-
providers/google/src/airflow/providers/google/cloud/operators/dataproc.py::DataprocSubmitJobOperator
1716
providers/google/src/airflow/providers/google/cloud/operators/functions.py::CloudFunctionDeployFunctionOperator
1817
providers/google/src/airflow/providers/google/cloud/operators/gcs.py::GCSFileTransformOperator
1918
providers/google/src/airflow/providers/google/cloud/sensors/bigquery_dts.py::BigQueryDataTransferServiceTransferRunSensor

0 commit comments

Comments
 (0)