This is an automated email from the ASF dual-hosted git repository.
potiuk pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git
The following commit(s) were added to refs/heads/main by this push:
new eb89b4092b9 Fix SSH commands failing when the task holds over 1024
file descriptors (#74211)
eb89b4092b9 is described below
commit eb89b4092b9da3c3495b65e780f3c2c4f953b267
Author: Firas Bouzazi <[email protected]>
AuthorDate: Mon Oct 5 17:57:14 2026 +0100
Fix SSH commands failing when the task holds over 1024 file descriptors
(#74211)
* Fix SSH commands failing when the task holds over 1024 file descriptors
select.select() cannot watch a descriptor numbered 1024 or higher, so
SSHHook.exec_ssh_client_command failed with "filedescriptor out of range
in select()" once the task process had that many descriptors open,
whatever its open-file limit. Task SDK workers reach that count far more
easily than Airflow 2 task processes did.
closes: #74205
* Pin the test SSH server's host key instead of accepting any key
The regression test generates the server's key itself, so the client can
trust exactly that key and keep paramiko's default rejection of unknown
host keys rather than accepting whatever key it is offered.
* Use an ECDSA host key for the in-process SSH test server
RSA is considered outdated, and ECDSA keys are also the fastest paramiko
can generate.
---
.../ssh/src/airflow/providers/ssh/hooks/ssh.py | 68 +++++++------
providers/ssh/tests/unit/ssh/hooks/test_ssh.py | 106 ++++++++++++++++++++-
2 files changed, 139 insertions(+), 35 deletions(-)
diff --git a/providers/ssh/src/airflow/providers/ssh/hooks/ssh.py
b/providers/ssh/src/airflow/providers/ssh/hooks/ssh.py
index 8674ce209fa..38f68d43131 100644
--- a/providers/ssh/src/airflow/providers/ssh/hooks/ssh.py
+++ b/providers/ssh/src/airflow/providers/ssh/hooks/ssh.py
@@ -20,11 +20,11 @@
from __future__ import annotations
import os
+import selectors
from base64 import decodebytes
from collections.abc import Sequence
from functools import cached_property
from io import StringIO
-from select import select
from typing import Any
import paramiko
@@ -501,36 +501,42 @@ class SSHHook(BaseHook):
timedout = False
- # read from both stdout and stderr
- while not channel.closed or channel.recv_ready() or
channel.recv_stderr_ready():
- readq, _, _ = select([channel], [], [], cmd_timeout)
- if cmd_timeout is not None:
- timedout = not readq
- for recv in readq:
- if recv.recv_ready():
- output = stdout.channel.recv(len(recv.in_buffer))
- agg_stdout += output
- for line in output.decode("utf-8",
"replace").strip("\n").splitlines():
- self.log.info(line)
- if recv.recv_stderr_ready():
- output =
stderr.channel.recv_stderr(len(recv.in_stderr_buffer))
- agg_stderr += output
- for line in output.decode("utf-8",
"replace").strip("\n").splitlines():
- self.log.warning(line)
- if (
- stdout.channel.exit_status_ready()
- and not stderr.channel.recv_stderr_ready()
- and not stdout.channel.recv_ready()
- ) or timedout:
- stdout.channel.shutdown_read()
- try:
- stdout.channel.close()
- except Exception:
- # there is a race that when shutdown_read has been called
and when
- # you try to close the connection, the socket is already
closed
- # We should ignore such errors (but we should log them
with warning)
- self.log.warning("Ignoring exception on close",
exc_info=True)
- break
+ # select.select() rejects descriptors numbered FD_SETSIZE (1024) or
above, which a task
+ # process can reach; DefaultSelector uses epoll/kqueue/poll where
available.
+ with selectors.DefaultSelector() as selector:
+ selector.register(channel, selectors.EVENT_READ, data=channel)
+
+ # read from both stdout and stderr
+ while not channel.closed or channel.recv_ready() or
channel.recv_stderr_ready():
+ events = selector.select(cmd_timeout)
+ if cmd_timeout is not None:
+ timedout = not events
+ for key, _ in events:
+ recv = key.data
+ if recv.recv_ready():
+ output = stdout.channel.recv(len(recv.in_buffer))
+ agg_stdout += output
+ for line in output.decode("utf-8",
"replace").strip("\n").splitlines():
+ self.log.info(line)
+ if recv.recv_stderr_ready():
+ output =
stderr.channel.recv_stderr(len(recv.in_stderr_buffer))
+ agg_stderr += output
+ for line in output.decode("utf-8",
"replace").strip("\n").splitlines():
+ self.log.warning(line)
+ if (
+ stdout.channel.exit_status_ready()
+ and not stderr.channel.recv_stderr_ready()
+ and not stdout.channel.recv_ready()
+ ) or timedout:
+ stdout.channel.shutdown_read()
+ try:
+ stdout.channel.close()
+ except Exception:
+ # there is a race that when shutdown_read has been
called and when
+ # you try to close the connection, the socket is
already closed
+ # We should ignore such errors (but we should log them
with warning)
+ self.log.warning("Ignoring exception on close",
exc_info=True)
+ break
stdout.close()
stderr.close()
diff --git a/providers/ssh/tests/unit/ssh/hooks/test_ssh.py
b/providers/ssh/tests/unit/ssh/hooks/test_ssh.py
index 3a6cf30528b..bc8fa2ae4eb 100644
--- a/providers/ssh/tests/unit/ssh/hooks/test_ssh.py
+++ b/providers/ssh/tests/unit/ssh/hooks/test_ssh.py
@@ -18,9 +18,14 @@
from __future__ import annotations
import json
+import os
import random
+import resource
+import selectors
+import socket
import string
import textwrap
+import threading
from io import StringIO
from unittest import mock
@@ -103,6 +108,87 @@ TEST_DISABLED_ALGORITHMS = {"pubkeys": ["rsa-sha2-256",
"rsa-sha2-512"]}
TEST_CIPHERS = ["aes128-ctr", "aes192-ctr", "aes256-ctr"]
+class _ExecServer(paramiko.ServerInterface):
+ """Answers every exec request with stdout, stderr and exit status 3, then
closes the channel."""
+
+ def get_allowed_auths(self, username):
+ return "password"
+
+ def check_auth_password(self, username, password):
+ return paramiko.AUTH_SUCCESSFUL
+
+ def check_channel_request(self, kind, chanid):
+ return paramiko.OPEN_SUCCEEDED
+
+ def check_channel_exec_request(self, channel, command):
+ def respond():
+ # Give the client time to see the exec request succeed before the
channel closes.
+ threading.Event().wait(0.2)
+ channel.sendall(b"out-1\n")
+ channel.sendall_stderr(b"err-1\n")
+ channel.sendall(b"out-2\n")
+ channel.send_exit_status(3)
+ channel.close()
+
+ threading.Thread(target=respond, daemon=True).start()
+ return True
+
+
[email protected]
+def in_process_ssh_client():
+ """Yield an SSH client connected to an in-process paramiko server over a
loopback socket."""
+ listener = socket.socket()
+ listener.bind(("127.0.0.1", 0))
+ listener.listen(1)
+ port = listener.getsockname()[1]
+ host_key = paramiko.ECDSAKey.generate()
+ transports = []
+
+ def serve():
+ sock, _ = listener.accept()
+ transport = paramiko.Transport(sock)
+ transport.add_server_key(host_key)
+ transport.start_server(server=_ExecServer())
+ transports.append(transport)
+
+ server_thread = threading.Thread(target=serve, daemon=True)
+ server_thread.start()
+ client = paramiko.SSHClient()
+ # Trust exactly the server's key; any other key is rejected by the default
policy.
+ client.get_host_keys().add(f"[127.0.0.1]:{port}", host_key.get_name(),
host_key)
+ client.connect(
+ "127.0.0.1",
+ port=port,
+ username="user",
+ password="password",
+ look_for_keys=False,
+ allow_agent=False,
+ )
+ yield client
+ client.close()
+ server_thread.join(timeout=10)
+ for transport in transports:
+ transport.close()
+ listener.close()
+
+
[email protected]
+def over_1024_open_fds():
+ """Hold enough descriptors that new ones are numbered above select()'s
FD_SETSIZE of 1024."""
+ count = 1100
+ soft, hard = resource.getrlimit(resource.RLIMIT_NOFILE)
+ if soft < count + 256:
+ if hard != resource.RLIM_INFINITY and hard < count + 256:
+ pytest.skip(f"RLIMIT_NOFILE hard limit {hard} is too low to open
{count} descriptors")
+ resource.setrlimit(resource.RLIMIT_NOFILE, (count + 256, hard))
+ fds = [os.open(os.devnull, os.O_RDONLY) for _ in range(count)]
+ assert max(fds) > 1024
+ yield
+ for fd in fds:
+ os.close(fd)
+ resource.setrlimit(resource.RLIMIT_NOFILE, (soft, hard))
+
+
class TestSSHHook:
CONN_SSH_WITH_NO_EXTRA = "ssh_with_no_extra"
CONN_SSH_WITH_PRIVATE_KEY_EXTRA = "ssh_with_private_key_extra"
@@ -953,14 +1039,18 @@ class TestSSHHook:
mock_client = mock.MagicMock(spec=paramiko.SSHClient)
mock_client.exec_command.return_value = (mock_stdin, mock_stdout,
mock_stderr)
- def fake_select(rlist, wlist, xlist, timeout=None):
- assert timeout == pytest.approx(0.001), f"Expected cmd_timeout
passed to select, got {timeout}"
- return [], [], []
+ mock_selector = mock.create_autospec(selectors.BaseSelector,
instance=True)
+ mock_selector.__enter__.return_value = mock_selector
+ mock_selector.select.return_value = []
- with mock.patch("airflow.providers.ssh.hooks.ssh.select",
side_effect=fake_select):
+ with mock.patch(
+ "airflow.providers.ssh.hooks.ssh.selectors.DefaultSelector",
return_value=mock_selector
+ ):
with pytest.raises(AirflowException, match="SSH command timed
out"):
hook.exec_ssh_client_command(mock_client, "sleep 1", False,
None)
+ assert mock_selector.select.call_args_list ==
[mock.call(pytest.approx(0.001))]
+
assert mock_client.exec_command.call_args_list == [
mock.call(command="sleep 1", get_pty=False, timeout=0.001,
environment=None)
]
@@ -971,6 +1061,14 @@ class TestSSHHook:
assert mock.call() in mock_stdout.close.call_args_list
assert mock.call() in mock_stderr.close.call_args_list
+ @pytest.mark.usefixtures("over_1024_open_fds")
+ def test_exec_ssh_client_command_with_descriptors_above_fd_setsize(self,
in_process_ssh_client):
+ hook = SSHHook(remote_host="localhost", cmd_timeout=10)
+
+ ret = hook.exec_ssh_client_command(in_process_ssh_client, "anything",
False, None)
+
+ assert ret == (3, b"out-1\nout-2\n", b"err-1\n")
+
def test_command_timeout_not_set(self, monkeypatch):
hook = SSHHook(
ssh_conn_id="ssh_default",