Skip to content

Commit 8a963a6

Browse files
committed
fix(providers/google): fallback to fetch Dataflow job ID by name prefix when Beam regex matches none
1 parent 067e33a commit 8a963a6

3 files changed

Lines changed: 93 additions & 0 deletions

File tree

  • providers

providers/apache/beam/src/airflow/providers/apache/beam/operators/beam.py

Lines changed: 25 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -460,6 +460,15 @@ def execute_on_dataflow(self, context: Context):
460460
)
461461

462462
location = self.dataflow_config.location or DEFAULT_DATAFLOW_LOCATION
463+
if not self.dataflow_job_id and self.dataflow_hook and self.dataflow_job_name:
464+
fetched_job_id = self.dataflow_hook.fetch_job_id_by_name(
465+
job_name=self.dataflow_job_name,
466+
project_id=self.dataflow_config.project_id,
467+
location=location,
468+
)
469+
if fetched_job_id and isinstance(fetched_job_id, str):
470+
self.dataflow_job_id = fetched_job_id
471+
463472
DataflowJobLink.persist(
464473
context=context,
465474
region=self.dataflow_config.location,
@@ -655,6 +664,14 @@ def execute_on_dataflow(self, context: Context):
655664
is_dataflow_job_id_exist_callback=self.is_dataflow_job_id_exist_callback,
656665
)
657666
if self.dataflow_job_name and self.dataflow_config.location:
667+
if not self.dataflow_job_id and self.dataflow_hook:
668+
fetched_job_id = self.dataflow_hook.fetch_job_id_by_name(
669+
job_name=self.dataflow_job_name,
670+
project_id=self.dataflow_config.project_id,
671+
location=self.dataflow_config.location,
672+
)
673+
if fetched_job_id and isinstance(fetched_job_id, str):
674+
self.dataflow_job_id = fetched_job_id
658675
DataflowJobLink.persist(
659676
context=context,
660677
region=self.dataflow_config.location,
@@ -823,6 +840,14 @@ def execute(self, context: Context):
823840
variables=snake_case_pipeline_options,
824841
process_line_callback=process_line_callback,
825842
)
843+
if not self.dataflow_job_id and dataflow_job_name and self.dataflow_config.location:
844+
fetched_job_id = self.dataflow_hook.fetch_job_id_by_name(
845+
job_name=dataflow_job_name,
846+
project_id=self.dataflow_config.project_id,
847+
location=self.dataflow_config.location,
848+
)
849+
if fetched_job_id and isinstance(fetched_job_id, str):
850+
self.dataflow_job_id = fetched_job_id
826851
DataflowJobLink.persist(context=context)
827852
if dataflow_job_name and self.dataflow_config.location:
828853
self.dataflow_hook.wait_for_done(

providers/google/src/airflow/providers/google/cloud/hooks/dataflow.py

Lines changed: 32 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1173,6 +1173,38 @@ def get_job(
11731173
)
11741174
return jobs_controller.fetch_job_by_id(job_id)
11751175

1176+
@GoogleBaseHook.fallback_to_default_project_id
1177+
def fetch_job_id_by_name(
1178+
self,
1179+
job_name: str,
1180+
project_id: str = PROVIDE_PROJECT_ID,
1181+
location: str = DEFAULT_DATAFLOW_LOCATION,
1182+
) -> str | None:
1183+
"""
1184+
Fetch the job ID of the job with the specified name prefix.
1185+
1186+
:param job_name: Job name prefix to search for.
1187+
:param project_id: Optional, the Google Cloud project ID.
1188+
:param location: The location of the Dataflow job.
1189+
:return: the Job ID if exactly one job is found, otherwise None.
1190+
"""
1191+
try:
1192+
jobs_controller = _DataflowJobsController(
1193+
dataflow=self.get_conn(),
1194+
project_number=project_id,
1195+
location=location,
1196+
)
1197+
jobs = jobs_controller._fetch_jobs_by_prefix_name(job_name)
1198+
if len(jobs) == 1:
1199+
return jobs[0]["id"]
1200+
if len(jobs) > 1:
1201+
self.log.warning("Multiple Dataflow jobs found matching prefix %s: %s", job_name, [j["name"] for j in jobs])
1202+
else:
1203+
self.log.info("No Dataflow jobs found matching prefix %s", job_name)
1204+
except Exception as e:
1205+
self.log.warning("Failed to fetch Dataflow job ID by name prefix %s: %s", job_name, e)
1206+
return None
1207+
11761208
@GoogleBaseHook.fallback_to_default_project_id
11771209
def fetch_job_metrics_by_id(
11781210
self,

providers/google/tests/unit/google/cloud/hooks/test_dataflow.py

Lines changed: 36 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -342,6 +342,42 @@ def test_wait_for_done(self, mock_conn, mock_dataflowjob):
342342
)
343343
method_wait_for_done.assert_called_once_with()
344344

345+
@mock.patch(DATAFLOW_STRING.format("_DataflowJobsController"))
346+
@mock.patch(DATAFLOW_STRING.format("DataflowHook.get_conn"))
347+
def test_fetch_job_id_by_name(self, mock_conn, mock_dataflowjob):
348+
controller = mock_dataflowjob.return_value
349+
controller._fetch_jobs_by_prefix_name.return_value = [{"id": "TEST_JOB_ID_123"}]
350+
351+
res = self.dataflow_hook.fetch_job_id_by_name(
352+
job_name="JOB_NAME",
353+
project_id=TEST_PROJECT_ID,
354+
location=TEST_LOCATION,
355+
)
356+
mock_conn.assert_called_once()
357+
mock_dataflowjob.assert_called_once_with(
358+
dataflow=mock_conn.return_value,
359+
project_number=TEST_PROJECT_ID,
360+
location=TEST_LOCATION,
361+
)
362+
controller._fetch_jobs_by_prefix_name.assert_called_once_with("JOB_NAME")
363+
assert res == "TEST_JOB_ID_123"
364+
365+
@mock.patch(DATAFLOW_STRING.format("_DataflowJobsController"))
366+
@mock.patch(DATAFLOW_STRING.format("DataflowHook.get_conn"))
367+
def test_fetch_job_id_by_name_multiple_jobs(self, mock_conn, mock_dataflowjob):
368+
controller = mock_dataflowjob.return_value
369+
controller._fetch_jobs_by_prefix_name.return_value = [
370+
{"id": "TEST_JOB_ID_123", "name": "JOB_NAME_1"},
371+
{"id": "TEST_JOB_ID_456", "name": "JOB_NAME_2"},
372+
]
373+
374+
res = self.dataflow_hook.fetch_job_id_by_name(
375+
job_name="JOB_NAME",
376+
project_id=TEST_PROJECT_ID,
377+
location=TEST_LOCATION,
378+
)
379+
assert res is None
380+
345381

346382
@pytest.mark.db_test
347383
class TestDataflowTemplateHook:

0 commit comments

Comments
 (0)