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 dbee3c6c11b Fix concurrent SFTP directory transfers dropping parent
hook's connection overrides (#73647)
dbee3c6c11b is described below
commit dbee3c6c11bce3401a0e3d5c8633bcc168dec5be
Author: abhishekmauryaKsolves <[email protected]>
AuthorDate: Mon Oct 5 19:03:32 2026 +0530
Fix concurrent SFTP directory transfers dropping parent hook's connection
overrides (#73647)
* fix(sftp): propagate parent hook connection overrides to
concurrent-transfer workers
SFTPHook.store_directory_concurrently() and
retrieve_directory_concurrently()
built each worker hook as SFTPHook(ssh_conn_id=self.ssh_conn_id), discarding
every other constructor override on the parent hook (remote_host, port,
username, password, key_file, proxy settings, timeouts, no_host_key_check,
etc.). Workers fell back to the connection's raw defaults instead of the
parent hook's effective, already-resolved settings.
Add SFTPHook._build_worker_hook(), which builds a worker hook from the
parent's effective connection settings, and use it in both
store_directory_concurrently() and retrieve_directory_concurrently().
Closes #73585
* Address review feedback: fix key_file/pkey guard, tighten tests
- Re-resolve key_file when pkey is set on the worker hook, avoiding the
key_file/private_key guard in SSHHook that could make worker construction
raise even though the parent hook succeeded.
- Use MagicMock(spec=Connection) in
test_build_worker_hook_inherits_parent_overrides
so the mock stays honest against the real Connection interface.
- Simplify
test_store_and_retrieve_directory_concurrently_use_parent_overrides
to a wiring test (assert _build_worker_hook is called once per worker);
value-level assertions are already covered by
test_build_worker_hook_inherits_parent_overrides.
---------
Co-authored-by: abhishekmauryaKsolves <[email protected]>
---
.../sftp/src/airflow/providers/sftp/hooks/sftp.py | 44 ++++++++++-
providers/sftp/tests/unit/sftp/hooks/test_sftp.py | 92 ++++++++++++++++++++++
2 files changed, 134 insertions(+), 2 deletions(-)
diff --git a/providers/sftp/src/airflow/providers/sftp/hooks/sftp.py
b/providers/sftp/src/airflow/providers/sftp/hooks/sftp.py
index bab90ca031d..96df4feb730 100644
--- a/providers/sftp/src/airflow/providers/sftp/hooks/sftp.py
+++ b/providers/sftp/src/airflow/providers/sftp/hooks/sftp.py
@@ -207,6 +207,46 @@ class SFTPHook(SSHHook):
"""Get the number of open connections."""
return self._conn_count
+ def _build_worker_hook(self) -> SFTPHook:
+ """
+ Build a new SFTPHook for a concurrent-transfer worker.
+
+ Mirrors this hook's effective connection settings -- i.e. the result
of merging
+ this hook's constructor overrides (``remote_host``, ``port``,
``username``, etc.)
+ with the underlying Airflow connection -- so worker hooks used by
+ ``store_directory_concurrently`` and
``retrieve_directory_concurrently`` connect
+ the same way the parent hook does, instead of falling back to the
connection's
+ raw defaults.
+ """
+ worker_hook = SFTPHook(
+ ssh_conn_id=self.ssh_conn_id,
+ remote_host=self.remote_host,
+ username=self.username,
+ password=self.password,
+ # Re-resolve key_file when pkey is set to avoid the
key_file/private_key guard.
+ key_file=None if self.pkey else self.key_file,
+ port=self.port,
+ conn_timeout=self.conn_timeout,
+ cmd_timeout=self.cmd_timeout,
+ keepalive_interval=self.keepalive_interval,
+ banner_timeout=self.banner_timeout,
+ disabled_algorithms=self.disabled_algorithms,
+ ciphers=self.ciphers,
+ auth_timeout=self.auth_timeout,
+ host_proxy_cmd=self.host_proxy_cmd,
+ conn_retry_attempts=self.conn_retry_attempts,
+ )
+ # These have no constructor parameter and are only ever resolved from
the
+ # connection's `extra` field or left at their class default, so copy
the
+ # parent's already-resolved values across explicitly.
+ worker_hook.no_host_key_check = self.no_host_key_check
+ worker_hook.allow_host_key_change = self.allow_host_key_change
+ worker_hook.host_key = self.host_key
+ worker_hook.look_for_keys = self.look_for_keys
+ worker_hook.compress = self.compress
+ worker_hook.pkey = self.pkey
+ return worker_hook
+
@handle_connection_management
def describe_directory(self, path: str) -> dict[str, dict[str, str | int |
None]]:
"""
@@ -478,7 +518,7 @@ class SFTPHook(SSHHook):
remote_file_chunks = [remote_file_paths[i::workers] for i in
range(workers)]
local_file_chunks = [new_local_file_paths[i::workers] for i in
range(workers)]
self.log.info("Opening %s new SFTP connections", workers)
- conns = [SFTPHook(ssh_conn_id=self.ssh_conn_id).get_conn() for _ in
range(workers)]
+ conns = [self._build_worker_hook().get_conn() for _ in range(workers)]
try:
self.log.info("Retrieving files concurrently with %s threads",
workers)
with concurrent.futures.ThreadPoolExecutor(max_workers=workers) as
executor:
@@ -571,7 +611,7 @@ class SFTPHook(SSHHook):
remote_file_chunks = [new_remote_file_paths[i::workers] for i in
range(workers)]
local_file_chunks = [local_file_paths[i::workers] for i in
range(workers)]
self.log.info("Opening %s new SFTP connections", workers)
- conns = [SFTPHook(ssh_conn_id=self.ssh_conn_id).get_conn() for _ in
range(workers)]
+ conns = [self._build_worker_hook().get_conn() for _ in range(workers)]
try:
self.log.info("Storing files concurrently with %s threads",
workers)
with concurrent.futures.ThreadPoolExecutor(max_workers=workers) as
executor:
diff --git a/providers/sftp/tests/unit/sftp/hooks/test_sftp.py
b/providers/sftp/tests/unit/sftp/hooks/test_sftp.py
index cf074a4a30c..c6cdf214915 100644
--- a/providers/sftp/tests/unit/sftp/hooks/test_sftp.py
+++ b/providers/sftp/tests/unit/sftp/hooks/test_sftp.py
@@ -634,6 +634,98 @@ class TestSFTPHook:
)
assert retrieved_dir_name in os.listdir(os.path.join(self.temp_dir,
TMP_DIR_FOR_TESTS))
+ @patch("airflow.providers.sftp.hooks.sftp.SFTPHook.get_connection")
+ def test_build_worker_hook_inherits_parent_overrides(self,
mock_get_connection):
+ """
+ Regression test for #73585.
+
+ SFTPHook._build_worker_hook() must copy the parent hook's *effective*
+ connection settings (constructor overrides merged with the connection)
+ onto the worker hook it builds for concurrent transfers, not just
+ ssh_conn_id / no_host_key_check.
+ """
+ mock_connection = MagicMock(spec=Connection)
+ mock_connection.login = "conn_user"
+ mock_connection.password = "conn_pass"
+ mock_connection.host = "conn.example.com"
+ mock_connection.port = 2222
+ mock_connection.extra = None
+ mock_get_connection.return_value = mock_connection
+
+ parent_hook = SFTPHook(
+ ssh_conn_id="sftp_default",
+ remote_host="override.example.com",
+ port=2022,
+ username="override_user",
+ password="override_pass",
+ key_file="/tmp/override_key",
+ conn_timeout=42,
+ host_proxy_cmd="ncat --proxy proxy_host:1234 %h %p",
+ )
+ # Simulate values that only ever come from the connection's `extra`
+ # field (no constructor parameter exists for these on SSHHook).
+ parent_hook.no_host_key_check = False
+ parent_hook.allow_host_key_change = True
+ parent_hook.look_for_keys = False
+
+ worker_hook = parent_hook._build_worker_hook()
+
+ assert worker_hook is not parent_hook
+ assert worker_hook.remote_host == "override.example.com"
+ assert worker_hook.port == 2022
+ assert worker_hook.username == "override_user"
+ assert worker_hook.password == "override_pass"
+ assert worker_hook.key_file == "/tmp/override_key"
+ assert worker_hook.conn_timeout == 42
+ assert worker_hook.host_proxy_cmd == "ncat --proxy proxy_host:1234 %h
%p"
+ assert worker_hook.no_host_key_check is False
+ assert worker_hook.allow_host_key_change is True
+ assert worker_hook.look_for_keys is False
+
+ def
test_store_and_retrieve_directory_concurrently_use_parent_overrides(self):
+ """
+ Regression test for #73585.
+
+ store_directory_concurrently() and retrieve_directory_concurrently()
must build
+ every worker hook via self._build_worker_hook(), so each worker
inherits the
+ parent hook's effective remote_host/port/username instead of falling
back to
+ the connection's own defaults.
+ """
+ workers = 2
+ built_hooks = []
+ original_build = SFTPHook._build_worker_hook
+
+ def spy_build(hook_self):
+ worker_hook = original_build(hook_self)
+ built_hooks.append(worker_hook)
+ return worker_hook
+
+ with (
+ patch.object(SFTPHook, "_build_worker_hook", autospec=True,
side_effect=spy_build) as mock_build,
+ patch.object(SFTPHook, "get_conn", return_value=MagicMock()),
+ ):
+ stored_dir_name = "stored_dir_override"
+ self.hook.store_directory_concurrently(
+ remote_full_path=os.path.join(self.temp_dir,
TMP_DIR_FOR_TESTS, stored_dir_name),
+ local_full_path=os.path.join(self.temp_dir, TMP_DIR_FOR_TESTS,
SUB_DIR),
+ workers=workers,
+ )
+ # Value-level assertions (remote_host/port/username/etc.) are
covered by
+ # test_build_worker_hook_inherits_parent_overrides; this test only
checks
+ # that every worker hook is built via self._build_worker_hook().
+ assert mock_build.call_count == workers
+
+ built_hooks.clear()
+ mock_build.reset_mock()
+
+ retrieved_dir_name = "retrieved_dir_override"
+ self.hook.retrieve_directory_concurrently(
+ remote_full_path=os.path.join(self.temp_dir,
TMP_DIR_FOR_TESTS, stored_dir_name),
+ local_full_path=os.path.join(self.temp_dir, TMP_DIR_FOR_TESTS,
retrieved_dir_name),
+ workers=workers,
+ )
+ assert mock_build.call_count == workers
+
def test_validate_within_directory_rejects_escape(self):
base = os.path.join(self.temp_dir, "download")
with pytest.raises(ValueError, match="outside the destination
directory"):