diff --git a/providers/sftp/src/airflow/providers/sftp/hooks/sftp.py b/providers/sftp/src/airflow/providers/sftp/hooks/sftp.py index bab90ca031d40..627077320a12f 100644 --- a/providers/sftp/src/airflow/providers/sftp/hooks/sftp.py +++ b/providers/sftp/src/airflow/providers/sftp/hooks/sftp.py @@ -21,6 +21,7 @@ import asyncio import concurrent.futures +import copy import datetime import functools import inspect @@ -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]]: """ @@ -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: @@ -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: diff --git a/providers/sftp/tests/unit/sftp/hooks/test_sftp.py b/providers/sftp/tests/unit/sftp/hooks/test_sftp.py index cf074a4a30c12..49635e96cde7c 100644 --- a/providers/sftp/tests/unit/sftp/hooks/test_sftp.py +++ b/providers/sftp/tests/unit/sftp/hooks/test_sftp.py @@ -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")