Skip to content
Merged
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 @@ -469,19 +469,28 @@ async def _cleanup_tmp_dir(tmp_dir: str) -> None:

@staticmethod
async def _beam_version(py_interpreter: str) -> str:
version_script_cmd = shlex.join([py_interpreter, "-c", _APACHE_BEAM_VERSION_SCRIPT])
proc = await asyncio.create_subprocess_shell(
version_script_cmd,
stdout=asyncio.subprocess.PIPE,
stderr=asyncio.subprocess.PIPE,
)
stdout, stderr = await proc.communicate()
if proc.returncode != 0:
start_error: OSError | None = None
returncode: int | None
try:
proc = await asyncio.create_subprocess_exec(
py_interpreter,
"-c",
_APACHE_BEAM_VERSION_SCRIPT,
stdout=asyncio.subprocess.PIPE,
stderr=asyncio.subprocess.PIPE,
)
except OSError as e:
start_error = e
stdout, stderr, returncode = b"", str(e).encode(), 1
else:
stdout, stderr = await proc.communicate()
returncode = proc.returncode
if returncode != 0:
msg = (
f"Unable to retrieve Apache Beam version, return code {proc.returncode}."
f"Unable to retrieve Apache Beam version, return code {returncode}."
f"\nstdout: {stdout.decode()}\nstderr: {stderr.decode()}"
)
raise AirflowException(msg)
raise AirflowException(msg) from start_error
return stdout.decode().strip()

async def start_python_pipeline_async(
Expand Down Expand Up @@ -627,16 +636,12 @@ async def run_beam_command_async(
:param process_line_callback: Optional callback which can be used to process
stdout and stderr to detect job id
"""
cmd_str_representation = " ".join(shlex.quote(c) for c in cmd)
log.info("Running command: %s", cmd_str_representation)

# Creating a separate asynchronous process
process = await asyncio.create_subprocess_shell(
cmd_str_representation,
shell=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
close_fds=True,
log.info("Running command: %s", " ".join(shlex.quote(c) for c in cmd))

process = await asyncio.create_subprocess_exec(
*cmd,
stdout=asyncio.subprocess.PIPE,
stderr=asyncio.subprocess.PIPE,
cwd=working_directory,
)
# Waits for Apache Beam pipeline to complete.
Expand Down
42 changes: 42 additions & 0 deletions providers/apache/beam/tests/unit/apache/beam/hooks/test_beam.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@
import pytest

from airflow.providers.apache.beam.hooks.beam import (
_APACHE_BEAM_VERSION_SCRIPT,
BeamAsyncHook,
BeamHook,
beam_options_to_args,
Expand Down Expand Up @@ -479,6 +480,23 @@ async def test_beam_version_error(self):
with pytest.raises(AirflowException, match="Unable to retrieve Apache Beam version"):
await BeamAsyncHook._beam_version("python1")

@pytest.mark.asyncio
async def test_beam_version_invokes_interpreter_without_shell(self):
fake_proc = AsyncMock()
fake_proc.communicate = AsyncMock(return_value=(b"2.39.0\n", b""))
fake_proc.returncode = 0
interpreter = r"C:\Program Files\Python\python.exe"
with (
mock.patch("asyncio.create_subprocess_exec", new=AsyncMock(return_value=fake_proc)) as mock_exec,
mock.patch("asyncio.create_subprocess_shell", new=AsyncMock()) as mock_shell,
):
version = await BeamAsyncHook._beam_version(interpreter)

mock_shell.assert_not_called()
mock_exec.assert_awaited_once()
assert mock_exec.await_args.args == (interpreter, "-c", _APACHE_BEAM_VERSION_SCRIPT)
assert version == "2.39.0"

@pytest.mark.asyncio
@mock.patch("airflow.providers.apache.beam.hooks.beam.BeamAsyncHook.run_beam_command_async")
async def test_start_pipline_async(self, mock_runner):
Expand Down Expand Up @@ -690,3 +708,27 @@ async def test_start_java_pipeline_async(self, mock_start_pipeline, job_class, c
command_prefix=command_prefix,
process_line_callback=None,
)

@pytest.mark.asyncio
async def test_run_beam_command_async_uses_exec_with_argv(self):
hook = BeamAsyncHook(runner=DEFAULT_RUNNER)
fake_proc = AsyncMock()
fake_proc.stdout.readline = AsyncMock(return_value=b"")
fake_proc.stderr.readline = AsyncMock(return_value=b"")
fake_proc.wait = AsyncMock(return_value=0)
cmd = [
r"C:\Program Files\Python\python.exe",
r"C:\Program Files\pipelines\word count.py",
"--output=gs://test/output",
]
with (
mock.patch("asyncio.create_subprocess_exec", new=AsyncMock(return_value=fake_proc)) as mock_exec,
mock.patch("asyncio.create_subprocess_shell", new=AsyncMock()) as mock_shell,
):
return_code = await hook.run_beam_command_async(cmd=cmd, log=logging.getLogger("beam-test"))

mock_shell.assert_not_called()
mock_exec.assert_awaited_once()
assert mock_exec.await_args.args == tuple(cmd)
assert mock_exec.await_args.kwargs.get("shell") is None
assert return_code == 0