diff --git a/providers/apache/beam/src/airflow/providers/apache/beam/operators/beam.py b/providers/apache/beam/src/airflow/providers/apache/beam/operators/beam.py index e42f1e1d60a39..8c1ba6c27ed28 100644 --- a/providers/apache/beam/src/airflow/providers/apache/beam/operators/beam.py +++ b/providers/apache/beam/src/airflow/providers/apache/beam/operators/beam.py @@ -460,6 +460,15 @@ def execute_on_dataflow(self, context: Context): ) location = self.dataflow_config.location or DEFAULT_DATAFLOW_LOCATION + if not self.dataflow_job_id and self.dataflow_hook and self.dataflow_job_name: + fetched_job_id = self.dataflow_hook.fetch_job_id_by_name( + job_name=self.dataflow_job_name, + project_id=self.dataflow_config.project_id, + location=location, + ) + if fetched_job_id and isinstance(fetched_job_id, str): + self.dataflow_job_id = fetched_job_id + DataflowJobLink.persist( context=context, region=self.dataflow_config.location, @@ -655,6 +664,14 @@ def execute_on_dataflow(self, context: Context): is_dataflow_job_id_exist_callback=self.is_dataflow_job_id_exist_callback, ) if self.dataflow_job_name and self.dataflow_config.location: + if not self.dataflow_job_id and self.dataflow_hook: + fetched_job_id = self.dataflow_hook.fetch_job_id_by_name( + job_name=self.dataflow_job_name, + project_id=self.dataflow_config.project_id, + location=self.dataflow_config.location, + ) + if fetched_job_id and isinstance(fetched_job_id, str): + self.dataflow_job_id = fetched_job_id DataflowJobLink.persist( context=context, region=self.dataflow_config.location, @@ -823,6 +840,14 @@ def execute(self, context: Context): variables=snake_case_pipeline_options, process_line_callback=process_line_callback, ) + if not self.dataflow_job_id and dataflow_job_name and self.dataflow_config.location: + fetched_job_id = self.dataflow_hook.fetch_job_id_by_name( + job_name=dataflow_job_name, + project_id=self.dataflow_config.project_id, + location=self.dataflow_config.location, + ) + if fetched_job_id and isinstance(fetched_job_id, str): + self.dataflow_job_id = fetched_job_id DataflowJobLink.persist(context=context) if dataflow_job_name and self.dataflow_config.location: self.dataflow_hook.wait_for_done( diff --git a/providers/google/src/airflow/providers/google/cloud/hooks/dataflow.py b/providers/google/src/airflow/providers/google/cloud/hooks/dataflow.py index ea1af01c1ad50..a0468422c651c 100644 --- a/providers/google/src/airflow/providers/google/cloud/hooks/dataflow.py +++ b/providers/google/src/airflow/providers/google/cloud/hooks/dataflow.py @@ -1173,6 +1173,38 @@ def get_job( ) return jobs_controller.fetch_job_by_id(job_id) + @GoogleBaseHook.fallback_to_default_project_id + def fetch_job_id_by_name( + self, + job_name: str, + project_id: str = PROVIDE_PROJECT_ID, + location: str = DEFAULT_DATAFLOW_LOCATION, + ) -> str | None: + """ + Fetch the job ID of the job with the specified name prefix. + + :param job_name: Job name prefix to search for. + :param project_id: Optional, the Google Cloud project ID. + :param location: The location of the Dataflow job. + :return: the Job ID if exactly one job is found, otherwise None. + """ + try: + jobs_controller = _DataflowJobsController( + dataflow=self.get_conn(), + project_number=project_id, + location=location, + ) + jobs = jobs_controller._fetch_jobs_by_prefix_name(job_name) + if len(jobs) == 1: + return jobs[0]["id"] + if len(jobs) > 1: + self.log.warning("Multiple Dataflow jobs found matching prefix %s: %s", job_name, [j["name"] for j in jobs]) + else: + self.log.info("No Dataflow jobs found matching prefix %s", job_name) + except Exception as e: + self.log.warning("Failed to fetch Dataflow job ID by name prefix %s: %s", job_name, e) + return None + @GoogleBaseHook.fallback_to_default_project_id def fetch_job_metrics_by_id( self, diff --git a/providers/google/tests/unit/google/cloud/hooks/test_dataflow.py b/providers/google/tests/unit/google/cloud/hooks/test_dataflow.py index 28d06fde772ac..a72e384a071ec 100644 --- a/providers/google/tests/unit/google/cloud/hooks/test_dataflow.py +++ b/providers/google/tests/unit/google/cloud/hooks/test_dataflow.py @@ -342,6 +342,42 @@ def test_wait_for_done(self, mock_conn, mock_dataflowjob): ) method_wait_for_done.assert_called_once_with() + @mock.patch(DATAFLOW_STRING.format("_DataflowJobsController")) + @mock.patch(DATAFLOW_STRING.format("DataflowHook.get_conn")) + def test_fetch_job_id_by_name(self, mock_conn, mock_dataflowjob): + controller = mock_dataflowjob.return_value + controller._fetch_jobs_by_prefix_name.return_value = [{"id": "TEST_JOB_ID_123"}] + + res = self.dataflow_hook.fetch_job_id_by_name( + job_name="JOB_NAME", + project_id=TEST_PROJECT_ID, + location=TEST_LOCATION, + ) + mock_conn.assert_called_once() + mock_dataflowjob.assert_called_once_with( + dataflow=mock_conn.return_value, + project_number=TEST_PROJECT_ID, + location=TEST_LOCATION, + ) + controller._fetch_jobs_by_prefix_name.assert_called_once_with("JOB_NAME") + assert res == "TEST_JOB_ID_123" + + @mock.patch(DATAFLOW_STRING.format("_DataflowJobsController")) + @mock.patch(DATAFLOW_STRING.format("DataflowHook.get_conn")) + def test_fetch_job_id_by_name_multiple_jobs(self, mock_conn, mock_dataflowjob): + controller = mock_dataflowjob.return_value + controller._fetch_jobs_by_prefix_name.return_value = [ + {"id": "TEST_JOB_ID_123", "name": "JOB_NAME_1"}, + {"id": "TEST_JOB_ID_456", "name": "JOB_NAME_2"}, + ] + + res = self.dataflow_hook.fetch_job_id_by_name( + job_name="JOB_NAME", + project_id=TEST_PROJECT_ID, + location=TEST_LOCATION, + ) + assert res is None + @pytest.mark.db_test class TestDataflowTemplateHook: