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 @@ -24,7 +24,7 @@
import time
import warnings
from collections.abc import MutableSequence, Sequence
from dataclasses import dataclass
from dataclasses import dataclass, replace
from datetime import datetime, timedelta
from enum import Enum
from functools import cached_property
Expand Down Expand Up @@ -66,6 +66,7 @@
from airflow.triggers.base import StartTriggerArgs

if TYPE_CHECKING:
import jinja2
from google.api_core import operation
from google.api_core.retry_async import AsyncRetry
from google.protobuf.duration_pb2 import Duration
Expand Down Expand Up @@ -2004,9 +2005,13 @@ def __init__(
self.openlineage_inject_parent_job_info = openlineage_inject_parent_job_info
self.openlineage_inject_transport_info = openlineage_inject_transport_info

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

def execute(self, context: Context):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -2491,6 +2491,7 @@ def test_start_from_trigger_sets_start_trigger_args(self):
cancel_on_kill=False,
request_id=REQUEST_ID,
)
op.render_template_fields({})
assert op.start_from_trigger is True
assert (
op.start_trigger_args.trigger_cls
Expand All @@ -2508,6 +2509,46 @@ def test_start_from_trigger_sets_start_trigger_args(self):
}
assert op.start_trigger_args.next_method == "execute_complete"

def test_start_trigger_args_use_rendered_template_fields(self):
op = DataprocSubmitJobOperator(
task_id=TASK_ID,
region="{{ params.region }}",
project_id="{{ params.project }}",
job={"reference": {"job_id": "job-{{ ds }}"}},
deferrable=True,
start_from_trigger=True,
)
op.render_template_fields(
{"ds": "2026-01-01", "params": {"region": GCP_REGION, "project": GCP_PROJECT}}
)
trigger_kwargs = op.start_trigger_args.trigger_kwargs
assert trigger_kwargs["region"] == GCP_REGION
assert trigger_kwargs["project_id"] == GCP_PROJECT
assert trigger_kwargs["job"] == {"reference": {"job_id": "job-2026-01-01"}}

def test_start_trigger_args_are_not_shared_between_tasks(self):
first = DataprocSubmitJobOperator(
task_id="first",
region=GCP_REGION,
project_id=GCP_PROJECT,
job={"reference": {"job_id": "first"}},
deferrable=True,
start_from_trigger=True,
)
second = DataprocSubmitJobOperator(
task_id="second",
region=GCP_REGION,
project_id=GCP_PROJECT,
job={"reference": {"job_id": "second"}},
deferrable=True,
start_from_trigger=True,
)
first.render_template_fields({})
second.render_template_fields({})
assert first.start_trigger_args.trigger_kwargs["job"] == {"reference": {"job_id": "first"}}
assert second.start_trigger_args.trigger_kwargs["job"] == {"reference": {"job_id": "second"}}
assert DataprocSubmitJobOperator.start_trigger_args.trigger_kwargs == {}

def test_start_from_trigger_without_deferrable_does_not_set_args(self):
op = DataprocSubmitJobOperator(
task_id=TASK_ID,
Expand All @@ -2518,6 +2559,7 @@ def test_start_from_trigger_without_deferrable_does_not_set_args(self):
deferrable=False,
start_from_trigger=True,
)
op.render_template_fields({})
assert op.start_from_trigger is True
assert op.start_trigger_args.trigger_kwargs == {}

Expand Down
1 change: 0 additions & 1 deletion scripts/ci/prek/validate_operators_init_exemptions.txt
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,6 @@ providers/google/src/airflow/providers/google/cloud/operators/cloud_batch.py::Cl
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
providers/google/src/airflow/providers/google/cloud/operators/dataproc.py::DataprocSubmitJobOperator
providers/google/src/airflow/providers/google/cloud/operators/functions.py::CloudFunctionDeployFunctionOperator
providers/google/src/airflow/providers/google/cloud/operators/gcs.py::GCSFileTransformOperator
providers/google/src/airflow/providers/google/cloud/sensors/bigquery_dts.py::BigQueryDataTransferServiceTransferRunSensor
Expand Down
Loading