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 @@ -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,
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
36 changes: 36 additions & 0 deletions providers/google/tests/unit/google/cloud/hooks/test_dataflow.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down