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",

Reply via email to