Skip to content
Open
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
24 changes: 22 additions & 2 deletions providers/sftp/src/airflow/providers/sftp/hooks/sftp.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@

import asyncio
import concurrent.futures
import copy
import datetime
import functools
import inspect
Expand Down Expand Up @@ -207,6 +208,25 @@ def get_conn_count(self) -> int:
"""Get the number of open connections."""
return self._conn_count

def _make_worker_hook(self) -> SFTPHook:
"""
Return a hook for one concurrent transfer worker.

Rebuilding the worker from ``ssh_conn_id`` would drop everything the caller passed to the
constructor (``remote_host``, port, credentials, proxy, host key settings), so a worker
could connect to a different host than the one the directory was listed on. A copy keeps
those settings; only the connection state is reset, so each worker opens its own
connection and proxy.
"""
worker = copy.copy(self)
worker.conn = None
worker.client = None
worker._ssh_conn = None
worker._sftp_conn = None
worker._conn_count = 0
worker.__dict__.pop("host_proxy", None)
return worker

@handle_connection_management
def describe_directory(self, path: str) -> dict[str, dict[str, str | int | None]]:
"""
Expand Down Expand Up @@ -478,7 +498,7 @@ def retrieve_file_chunk(
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._make_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:
Expand Down Expand Up @@ -571,7 +591,7 @@ def store_file_chunk(
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._make_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:
Expand Down
49 changes: 49 additions & 0 deletions providers/sftp/tests/unit/sftp/hooks/test_sftp.py
Original file line number Diff line number Diff line change
Expand Up @@ -352,6 +352,55 @@ def test_get_mod_time(self):
)
assert len(output) == 14

@patch("airflow.providers.sftp.hooks.sftp.SFTPHook.get_connection")
def test_make_worker_hook_keeps_effective_settings(self, get_connection):
get_connection.return_value = Connection(
login="conn-user", host="conn-host", port=22, extra=json.dumps({"no_host_key_check": "false"})
)
hook = SFTPHook(
remote_host="explicit-host", username="explicit-user", port=2222, host_proxy_cmd="nc %h %p"
)
hook.client = MagicMock(spec=SSHClient)
hook.conn = MagicMock(spec=SFTPClient)

worker = hook._make_worker_hook()

assert worker is not hook
assert (worker.remote_host, worker.username, worker.port, worker.host_proxy_cmd) == (
"explicit-host",
"explicit-user",
2222,
"nc %h %p",
)
assert worker.no_host_key_check is False
assert worker.client is None
assert worker.conn is None
assert worker.get_conn_count() == 0

@pytest.mark.parametrize("direction", ["retrieve", "store"])
@patch.object(SFTPHook, "_make_worker_hook", autospec=True)
@patch("airflow.providers.sftp.hooks.sftp.SFTPHook.get_connection")
def test_concurrent_transfers_use_worker_hooks(
self, get_connection, mock_make_worker_hook, direction, tmp_path
):
get_connection.return_value = Connection(login="login", host="host")
hook = SFTPHook(remote_host="explicit-host")
local_dir = tmp_path / "local"
with (
patch.object(SFTPHook, "get_managed_conn", autospec=True),
patch.object(SFTPHook, "get_tree_map", autospec=True, return_value=(["/remote/a.txt"], [], [])),
patch.object(SFTPHook, "path_exists", autospec=True, return_value=False),
patch.object(SFTPHook, "create_directory", autospec=True),
):
if direction == "retrieve":
hook.retrieve_directory_concurrently("/remote", str(local_dir), workers=2)
else:
local_dir.mkdir()
(local_dir / "a.txt").write_text("a")
hook.store_directory_concurrently("/remote", str(local_dir), workers=2)

assert mock_make_worker_hook.call_args_list == [call(hook), call(hook)]

@patch("airflow.providers.sftp.hooks.sftp.SFTPHook.get_connection")
def test_no_host_key_check_default(self, get_connection):
connection = Connection(login="login", host="host")
Expand Down