Skip to content
Draft
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 @@ -125,6 +125,7 @@ def process_fd(
log: Logger,
process_line_callback: Callable[[str], None] | None = None,
is_dataflow_job_id_exist_callback: Callable[[], bool] | None = None,
drain: bool = False,
):
"""
Print output to logs.
Expand All @@ -134,6 +135,7 @@ def process_fd(
:param process_line_callback: Optional callback which can be used to process
stdout and stderr to detect job id.
:param log: logger.
:param drain: Whether to read until EOF instead of returning after one line.
"""
if fd not in (proc.stdout, proc.stderr):
raise AirflowException("No data in stderr or in stdout.")
Expand All @@ -148,6 +150,8 @@ def process_fd(
func_log(line.rstrip("\n"))
if is_dataflow_job_id_exist_callback and is_dataflow_job_id_exist_callback():
return
if not drain:
return


def run_beam_command(
Expand Down Expand Up @@ -191,18 +195,31 @@ def run_beam_command(
if is_dataflow_job_id_exist_callback and is_dataflow_job_id_exist_callback():
return

if is_dataflow_job_id_exist_callback and is_dataflow_job_id_exist_callback():
return

if proc.poll() is not None:
break

# Corner case: check if more output was created between the last read and the process termination
for readable_fd in reads:
process_fd(proc, readable_fd, log, process_line_callback, is_dataflow_job_id_exist_callback)
process_fd(
proc,
readable_fd,
log,
process_line_callback,
is_dataflow_job_id_exist_callback,
drain=True,
)

log.info("Process exited with return code: %s", proc.returncode)

if proc.returncode != 0:
raise AirflowException(f"Apache Beam process failed with return code {proc.returncode}")

if is_dataflow_job_id_exist_callback:
is_dataflow_job_id_exist_callback()


class BeamHook(BaseHook):
"""
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@
import os
import stat
import tempfile
import time
from abc import ABC, ABCMeta, abstractmethod
from collections.abc import Callable, Sequence
from concurrent.futures import ThreadPoolExecutor, as_completed
Expand All @@ -45,6 +46,7 @@
if TYPE_CHECKING:
from airflow.providers.common.compat.sdk import Context
GOOGLE_PROVIDER = ProvidersManager().providers.get("apache-airflow-providers-google")
_DATAFLOW_JOB_ID_LOOKUP_INTERVAL: float = 5.0


if GOOGLE_PROVIDER:
Expand Down Expand Up @@ -102,7 +104,7 @@ def _set_dataflow(
pipeline_options, dataflow_job_name, job_name_variable_key
)
process_line_callback = self.__get_dataflow_process_callback()
is_dataflow_job_id_exist_callback = self.__is_dataflow_job_id_exist_callback()
is_dataflow_job_id_exist_callback = self.__get_dataflow_job_id_callback(dataflow_job_name)
return dataflow_job_name, pipeline_options, process_line_callback, is_dataflow_job_id_exist_callback

def __set_dataflow_hook(self) -> DataflowHook:
Expand Down Expand Up @@ -152,11 +154,33 @@ def set_current_dataflow_job_id(job_id):
on_new_job_id_callback=set_current_dataflow_job_id
)

def __is_dataflow_job_id_exist_callback(self) -> Callable[[], bool]:
def is_dataflow_job_id_exist() -> bool:
return True if self.dataflow_job_id else False
def __get_dataflow_job_id_callback(self, dataflow_job_name: str) -> Callable[[], bool]:
last_dataflow_job_id_lookup = -_DATAFLOW_JOB_ID_LOOKUP_INTERVAL

return is_dataflow_job_id_exist
def try_resolve_dataflow_job_id() -> bool:
nonlocal last_dataflow_job_id_lookup

if self.dataflow_job_id:
return True
if not self.dataflow_hook:
return False

now = time.monotonic()
if now - last_dataflow_job_id_lookup < _DATAFLOW_JOB_ID_LOOKUP_INTERVAL:
return False
last_dataflow_job_id_lookup = now

job_id = self.dataflow_hook.fetch_job_id_by_name(
prefix_name=dataflow_job_name,
project_id=self.dataflow_config.project_id,
location=self.dataflow_config.location or DEFAULT_DATAFLOW_LOCATION,
)
if job_id:
self.dataflow_job_id = job_id
return True
return False

return try_resolve_dataflow_job_id


class BeamBasePipelineOperator(BaseOperator, BeamDataflowMixin, ABC):
Expand Down Expand Up @@ -467,7 +491,7 @@ def execute_on_dataflow(self, context: Context):
project_id=self.dataflow_config.project_id,
)

if self.deferrable:
if self.deferrable and self.dataflow_job_id:
trigger_args = {
"job_id": self.dataflow_job_id,
"project_id": self.dataflow_config.project_id,
Expand Down Expand Up @@ -661,7 +685,7 @@ def execute_on_dataflow(self, context: Context):
job_id=self.dataflow_job_id,
project_id=self.dataflow_config.project_id,
)
if self.deferrable:
if self.deferrable and self.dataflow_job_id:
trigger_args = {
"job_id": self.dataflow_job_id,
"project_id": self.dataflow_config.project_id,
Expand Down
36 changes: 33 additions & 3 deletions providers/apache/beam/tests/unit/apache/beam/hooks/test_beam.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@
BeamAsyncHook,
BeamHook,
beam_options_to_args,
process_fd,
run_beam_command,
)
from airflow.providers.common.compat.sdk import AirflowException
Expand Down Expand Up @@ -408,12 +409,11 @@ def test_beam_wait_for_done_logging(self, mock_select, mock_popen, caplog):
fake_stderr_fd.readline.side_effect = [
b"apache-beam-stderr-1",
b"apache-beam-stderr-2",
StopIteration,
b"apache-beam-stderr-3",
StopIteration,
b"apache-beam-other-stderr",
b"",
]
fake_stdout_fd.readline.side_effect = [b"apache-beam-stdout", StopIteration]
fake_stdout_fd.readline.side_effect = [b"apache-beam-stdout", b""]
mock_select.side_effect = [
([fake_stderr_fd], None, None),
(None, None, None),
Expand All @@ -440,6 +440,36 @@ def test_beam_wait_for_done_logging(self, mock_select, mock_popen, caplog):
assert "apache-beam-stderr-3" in warn_messages
assert "apache-beam-other-stderr" in warn_messages

def test_process_fd_reads_one_line_without_draining(self):
fake_logger = logging.getLogger("fake-beam-process-fd-logger")
mock_proc = MagicMock(name="FakeProc")
mock_proc.stderr = MagicMock(name="FakeStderr")
mock_proc.stdout = MagicMock(name="FakeStdout")
mock_proc.stdout.readline.side_effect = [b"apache-beam-stdout-1", b"apache-beam-stdout-2"]

process_fd(mock_proc, mock_proc.stdout, fake_logger)

mock_proc.stdout.readline.assert_called_once_with()

@mock.patch("subprocess.Popen")
@mock.patch("select.select")
def test_beam_wait_for_done_checks_dataflow_job_id_after_select_timeout(self, mock_select, mock_popen):
fake_logger = logging.getLogger("fake-beam-wait-for-done-logger")
mock_proc = MagicMock(name="FakeProc")
mock_proc.stderr = MagicMock(name="FakeStderr")
mock_proc.stdout = MagicMock(name="FakeStdout")
mock_proc.poll.side_effect = AssertionError("poll should not run after job id is resolved")
mock_popen.return_value = mock_proc
mock_select.return_value = ([], None, None)
is_dataflow_job_id_exist_callback = MagicMock(return_value=True)

run_beam_command(
["fake", "cmd"], fake_logger, is_dataflow_job_id_exist_callback=is_dataflow_job_id_exist_callback
)

is_dataflow_job_id_exist_callback.assert_called_once_with()
mock_proc.poll.assert_not_called()


class TestBeamOptionsToArgs:
@pytest.mark.parametrize(
Expand Down
129 changes: 129 additions & 0 deletions providers/apache/beam/tests/unit/apache/beam/operators/test_beam.py
Original file line number Diff line number Diff line change
Expand Up @@ -1029,6 +1029,71 @@ def test_exec_dataflow_runner(self, gcs_hook_mock, dataflow_hook_mock, beam_hook
)
beam_hook_mock.return_value.start_python_pipeline.assert_called_once()

@mock.patch(BEAM_OPERATOR_PATH.format("BeamHook"))
@mock.patch(BEAM_OPERATOR_PATH.format("DataflowHook"))
@mock.patch(BEAM_OPERATOR_PATH.format("GCSHook"))
def test_exec_dataflow_runner_resolves_job_id_by_name_before_deferring(
self, gcs_hook_mock, dataflow_hook_mock, beam_hook_mock
):
dataflow_config = DataflowConfiguration(impersonation_chain=TEST_IMPERSONATION_ACCOUNT)
op = BeamRunPythonPipelineOperator(
runner="DataflowRunner",
dataflow_config=dataflow_config,
**self.default_op_kwargs,
)
dataflow_hook_mock.return_value.fetch_job_id_by_name.return_value = JOB_ID

def start_python_pipeline(**kwargs):
assert kwargs["is_dataflow_job_id_exist_callback"]() is True

beam_hook_mock.return_value.start_python_pipeline.side_effect = start_python_pipeline

with pytest.raises(TaskDeferred) as exc:
op.execute(context=mock.MagicMock())

assert op.dataflow_job_id == JOB_ID
assert exc.value.trigger.job_id == JOB_ID
dataflow_hook_mock.return_value.fetch_job_id_by_name.assert_called_once_with(
prefix_name=dataflow_hook_mock.build_dataflow_job_name.return_value,
project_id=dataflow_hook_mock.return_value.project_id,
location="us-central1",
)

beam_hook_mock.return_value.start_python_pipeline.assert_called_once()

@mock.patch(BEAM_OPERATOR_PATH.format("DataflowJobLink.persist"))
@mock.patch(BEAM_OPERATOR_PATH.format("BeamHook"))
@mock.patch(BEAM_OPERATOR_PATH.format("DataflowHook"))
@mock.patch(BEAM_OPERATOR_PATH.format("GCSHook"))
def test_exec_dataflow_runner_without_job_id_falls_back_to_sync_wait(
self, gcs_hook_mock, dataflow_hook_mock, beam_hook_mock, persist_link_mock
):
dataflow_config = DataflowConfiguration(impersonation_chain=TEST_IMPERSONATION_ACCOUNT)
op = BeamRunPythonPipelineOperator(
runner="DataflowRunner",
dataflow_config=dataflow_config,
**self.default_op_kwargs,
)
dataflow_hook_mock.return_value.fetch_job_id_by_name.return_value = None

def start_python_pipeline(**kwargs):
assert kwargs["is_dataflow_job_id_exist_callback"]() is False

beam_hook_mock.return_value.start_python_pipeline.side_effect = start_python_pipeline

result = op.execute(context=mock.MagicMock())

assert result == {"dataflow_job_id": None}
dataflow_hook_mock.return_value.wait_for_done.assert_called_once_with(
job_name=dataflow_hook_mock.build_dataflow_job_name.return_value,
location="us-central1",
job_id=None,
project_id=dataflow_hook_mock.return_value.project_id,
)
persist_link_mock.assert_called_once()

beam_hook_mock.return_value.start_python_pipeline.assert_called_once()

@mock.patch(BEAM_OPERATOR_PATH.format("DataflowJobLink.persist"))
@mock.patch(BEAM_OPERATOR_PATH.format("BeamHook"))
@mock.patch(BEAM_OPERATOR_PATH.format("GCSHook"))
Expand Down Expand Up @@ -1157,6 +1222,70 @@ def test_exec_dataflow_runner(self, gcs_hook_mock, dataflow_hook_mock, beam_hook
)
beam_hook_mock.return_value.start_python_pipeline.assert_not_called()

@mock.patch(BEAM_OPERATOR_PATH.format("BeamHook"))
@mock.patch(BEAM_OPERATOR_PATH.format("DataflowHook"))
@mock.patch(BEAM_OPERATOR_PATH.format("GCSHook"))
def test_exec_dataflow_runner_resolves_job_id_by_name_before_deferring(
self, gcs_hook_mock, dataflow_hook_mock, beam_hook_mock
):
dataflow_config = DataflowConfiguration(impersonation_chain=TEST_IMPERSONATION_ACCOUNT)
op = BeamRunJavaPipelineOperator(
runner="DataflowRunner", dataflow_config=dataflow_config, **self.default_op_kwargs
)
dataflow_hook_mock.return_value.is_job_dataflow_running.return_value = False
dataflow_hook_mock.return_value.fetch_job_id_by_name.return_value = JOB_ID

def start_java_pipeline(**kwargs):
assert kwargs["is_dataflow_job_id_exist_callback"]() is True

beam_hook_mock.return_value.start_java_pipeline.side_effect = start_java_pipeline

with pytest.raises(TaskDeferred) as exc:
op.execute(context=mock.MagicMock())

assert op.dataflow_job_id == JOB_ID
assert exc.value.trigger.job_id == JOB_ID
dataflow_hook_mock.return_value.fetch_job_id_by_name.assert_called_once_with(
prefix_name=dataflow_hook_mock.build_dataflow_job_name.return_value,
project_id=dataflow_hook_mock.return_value.project_id,
location="us-central1",
)

beam_hook_mock.return_value.start_python_pipeline.assert_not_called()

@mock.patch(BEAM_OPERATOR_PATH.format("DataflowJobLink.persist"))
@mock.patch(BEAM_OPERATOR_PATH.format("BeamHook"))
@mock.patch(BEAM_OPERATOR_PATH.format("DataflowHook"))
@mock.patch(BEAM_OPERATOR_PATH.format("GCSHook"))
def test_exec_dataflow_runner_without_job_id_falls_back_to_sync_wait(
self, gcs_hook_mock, dataflow_hook_mock, beam_hook_mock, persist_link_mock
):
dataflow_config = DataflowConfiguration(impersonation_chain=TEST_IMPERSONATION_ACCOUNT)
op = BeamRunJavaPipelineOperator(
runner="DataflowRunner", dataflow_config=dataflow_config, **self.default_op_kwargs
)
dataflow_hook_mock.return_value.is_job_dataflow_running.return_value = False
dataflow_hook_mock.return_value.fetch_job_id_by_name.return_value = None

def start_java_pipeline(**kwargs):
assert kwargs["is_dataflow_job_id_exist_callback"]() is False

beam_hook_mock.return_value.start_java_pipeline.side_effect = start_java_pipeline

result = op.execute(context=mock.MagicMock())

assert result == {"dataflow_job_id": None}
dataflow_hook_mock.return_value.wait_for_done.assert_called_once_with(
job_name=dataflow_hook_mock.build_dataflow_job_name.return_value,
location="us-central1",
job_id=None,
multiple_jobs=False,
project_id=dataflow_hook_mock.return_value.project_id,
)
persist_link_mock.assert_called_once()

beam_hook_mock.return_value.start_python_pipeline.assert_not_called()

@mock.patch(BEAM_OPERATOR_PATH.format("DataflowJobLink.persist"))
@mock.patch(BEAM_OPERATOR_PATH.format("BeamHook"))
@mock.patch(BEAM_OPERATOR_PATH.format("GCSHook"))
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -462,6 +462,12 @@ def _fetch_jobs_by_prefix_name(self, prefix_name: str) -> list[dict]:
jobs = [job for job in jobs if job["name"].startswith(prefix_name)]
return jobs

def fetch_job_id_by_name(self, prefix_name: str) -> str | None:
jobs = self._fetch_jobs_by_prefix_name(prefix_name)
if len(jobs) == 1:
return jobs[0]["id"]
return None

def _refresh_jobs(self) -> None:
"""
Get all jobs by name.
Expand Down Expand Up @@ -1173,6 +1179,30 @@ 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,
prefix_name: str,
project_id: str = PROVIDE_PROJECT_ID,
location: str = DEFAULT_DATAFLOW_LOCATION,
) -> str | None:
"""
Fetch the Dataflow job ID for a unique job name prefix.

:param prefix_name: Job name prefix to search for.
:param project_id: Optional, the Google Cloud project ID in which to start a job.
If set to None or missing, the default project_id from the Google Cloud connection is used.
:param location: The location of the Dataflow job (for example europe-west1). See:
https://cloud.google.com/dataflow/docs/concepts/regional-endpoints
:return: the job ID if exactly one matching job exists, otherwise None.
"""
jobs_controller = _DataflowJobsController(
dataflow=self.get_conn(),
project_number=project_id,
location=location,
)
return jobs_controller.fetch_job_id_by_name(prefix_name.lower())

@GoogleBaseHook.fallback_to_default_project_id
def fetch_job_metrics_by_id(
self,
Expand Down
Loading