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 2f5d44acbb1 Add deferrable mode to SFTPOperator (#72336)
2f5d44acbb1 is described below
commit 2f5d44acbb1117883c0c741b00575fbfc9c0f402
Author: David Blain <[email protected]>
AuthorDate: Mon Sep 21 16:29:29 2026 +0200
Add deferrable mode to SFTPOperator (#72336)
* Add deferrable mode to SFTPOperator
closes: #65475
Co-authored-by: Copilot <[email protected]>
* Fix pyproject.toml: restore missing flit/uv sections stripped by original
PR
* Fix circular imports, deferrable wiring and dependency issues
- Fix circular import: SFTPOperationTrigger now imported lazily inside
execute() so operators -> triggers -> hooks is a clean DAG with no cycle
- Fix triggers/sftp.py to import SFTPOperation from hooks (where it lives)
instead of operators, fully breaking the circular dependency
- Wire up self.defer() in execute() when deferrable=True — the call was
missing so deferrable mode never actually deferred
- Apply conf.getboolean('operators', 'default_deferrable') as the default
for the deferrable parameter, matching Airflow convention
- Remove paramiko<5.0.0 upper cap: paramiko 5.0.0 is already in uv.lock
and the cap caused pip check failures in the CI image build
- Remove asgiref dependency: it was added by the original PR but is not
imported anywhere in the provider source
Co-authored-by: Copilot <[email protected]>
* Apply validate_within_directory to block path traversal in
retrieve_directory and warn on missing DELETE target
* Replace new raise AirflowException in SFTPOperator.execute with dedicated
SFTPOperationError
* Make sure SFTPOperation is tested in same unit test module as hooks
* Fix SFTPTrigger timezone import to use provider compat shim
* Refactor SFTP triggers to share hook logic and simplify SFTPTrigger
SFTPTrigger and the operator's deferred transfer trigger duplicated
connection/hook setup and, in SFTPTrigger's case, embedded file-sensing
logic directly in the trigger's run loop. Move that sensing logic into
SFTPHookAsync (sense_files_by_pattern/sense_path) so triggers stay thin
dispatchers, and introduce BaseSFTPTrigger to share sftp_conn_id storage
and hook construction between SFTPTrigger and the renamed
SFTPTransferTrigger (formerly SFTPOperationTrigger).
* Cache SFTPTrigger.newer_than as a UTC-converted property
* Import conf from common compat provider in SFTPOperator
* Remove provider newsfragment for SFTP deferrable mode PR
Providers are released from main and their changelogs are regenerated
from git log by the release manager; per-PR newsfragments are not
consumed for provider distributions.
* Forward remote_host override to async hook in deferrable SFTPOperator
* Restore case-insensitive operation handling in SFTPOperator
* Add directory and delete support to deferrable SFTPOperator transfers
SFTPHookAsync.transfer() only handled single-file GET/PUT and a bare
unlink for DELETE, so deferrable=True silently diverged from the
synchronous SFTPHook.transfer() whenever a directory path was used —
directory GET/PUT was unsupported and DELETE on a directory would
fail. Bring the async path to parity with the sync hook: dispatch to
directory-aware transfers when the source/target is a directory,
warn instead of raising when deleting an already-missing file, honor
prefetch/confirm options, and validate that downloaded paths stay
within the destination directory.
Also drop the idle SSH connection SFTPHookAsync.transfer() kept open
for its whole duration; each helper now opens its own connection like
the rest of the class.
* Fix deferrable SFTPOperator ignoring a supplied sftp_hook's connection
When only sftp_hook (not ssh_conn_id) was provided, deferring silently
fell back to SFTPHookAsync's default connection instead of the hook's
actual connection id, redirecting the deferred transfer to the wrong
server with no warning and a task that still reports success.
* Fix docs spellcheck failure in async SFTP hook docstrings
The word "pipelined" is not in the docs spelling wordlist and broke the
documentation build for the SFTP provider. Rewording the prefetch
descriptions is preferable to whitelisting a non-dictionary word.
* Raise FileExistsError from async SFTP directory transfers
The async retrieve_directory and store_directory methods copied the
synchronous methods' AirflowException for an existing target path. The
project no longer accepts new direct AirflowException usages, and
FileExistsError names the actual condition.
* Fix mypy missing-return errors in async SFTP hook isdir and path_exists
mypy is right that control can fall out of the async with blocks when a
context manager suppresses an exception, which would return None from a
method typed as returning bool. Assigning the result inside the block and
returning it once after the connection closes makes the return path
explicit instead of silencing the check.
* Fix sftp_hook regression test failing at hook construction
SSHHook resolves its connection eagerly, so the test raised
AirflowNotFoundException before the operator under test was even built.
Defining the connection through the environment keeps the test db-free
and lets it exercise the deferred path it was written for.
* Stop os.walk from blocking the Triggerer event loop in async SFTP
directory uploads
store_directory walked the local tree inline in a coroutine, stalling every
other deferred task in the Triggerer for the duration of the walk, while the
pure string work around it was needlessly pushed to threads. Collect the
walk
in one thread hop and compute relative paths with pathlib, which does no I/O
and keeps the async lint rule satisfied honestly.
* Drop unused aiofiles dependency from the SFTP provider
The async hook's retrieve_file used aiofiles to stream a remote file into a
local path. It now delegates to asyncssh's own get, and nothing else in the
provider imports aiofiles, so the dependency only added install weight.
* Reuse one SFTP connection across async directory transfers
A deferred directory transfer opened a fresh SSH connection for the isdir
check, another to walk the tree, and one more per file, so a large directory
cost thousands of handshakes where the synchronous path needs one. Every
operation now has a private helper that works on an already-open SFTP
client,
and transfer() opens a single connection per top-level path and passes that
client down, matching the synchronous hook's connection cost.
* Fix mypy error on async SFTP retrieve_file local path argument
asyncssh types the destination of get() as PurePath rather than the wider
os.PathLike, so the narrowed str | PathLike value did not type-check.
os.fspath is the exact PathLike contract and leaves strings untouched.
* Fail early when deferrable SFTPOperator has no connection id
When neither ssh_conn_id nor the supplied sftp_hook carries a connection
id, the deferred path cannot recover: the trigger would silently receive
None. Raising in execute() surfaces the misconfiguration where it can be
fixed, and gives mypy the narrowing it needs without casting.
* Share one connection across async SFTP hook calls like the synchronous
hook
The async hook threaded an open SFTP client through private helpers to avoid
a connection per operation, duplicating every method. The synchronous hook
already solves this with a reference-counted managed connection and a
decorator, so the same decorator now wraps coroutine methods too. Nested
calls reuse the open connection and a whole deferred transfer runs over a
single one, which is the resource that matters in the Triggerer.
* Test that deleting a missing SFTP file warns for paramiko's error shape
The operator used to recognise a missing path through a helper that also
matched an OSError carrying ENOENT. paramiko raises IOError(errno.ENOENT,
text), which Python turns into FileNotFoundError, so catching that class is
sufficient. Covering that exact shape in the test locks the decision in.
* Update providers/sftp/src/airflow/providers/sftp/hooks/sftp.py
Co-authored-by: Jarek Potiuk <[email protected]>
---------
Co-authored-by: Copilot <[email protected]>
Co-authored-by: Jarek Potiuk <[email protected]>
---
providers/sftp/README.rst | 1 -
providers/sftp/docs/index.rst | 1 -
providers/sftp/pyproject.toml | 1 -
.../sftp/src/airflow/providers/sftp/exceptions.py | 4 +
.../sftp/src/airflow/providers/sftp/hooks/sftp.py | 722 +++++++++++++++++----
.../src/airflow/providers/sftp/operators/sftp.py | 222 +++----
.../src/airflow/providers/sftp/triggers/sftp.py | 143 +++-
providers/sftp/tests/unit/sftp/hooks/test_sftp.py | 466 ++++++++++++-
.../sftp/tests/unit/sftp/operators/test_sftp.py | 211 +++++-
.../sftp/tests/unit/sftp/triggers/test_sftp.py | 26 +-
uv.lock | 2 -
11 files changed, 1471 insertions(+), 328 deletions(-)
diff --git a/providers/sftp/README.rst b/providers/sftp/README.rst
index 6bf1787d126..72e17a4557f 100644
--- a/providers/sftp/README.rst
+++ b/providers/sftp/README.rst
@@ -53,7 +53,6 @@ Requirements
==========================================
======================================
PIP package Version required
==========================================
======================================
-``aiofiles`` ``>=23.2.0``
``apache-airflow`` ``>=2.11.0``
``apache-airflow-providers-ssh`` ``>=6.0.0``
``apache-airflow-providers-common-compat`` ``>=1.12.0``
diff --git a/providers/sftp/docs/index.rst b/providers/sftp/docs/index.rst
index bdd9129aba9..18816a404bb 100644
--- a/providers/sftp/docs/index.rst
+++ b/providers/sftp/docs/index.rst
@@ -100,7 +100,6 @@ The minimum Apache Airflow version supported by this
provider distribution is ``
==========================================
======================================
PIP package Version required
==========================================
======================================
-``aiofiles`` ``>=23.2.0``
``apache-airflow`` ``>=2.11.0``
``apache-airflow-providers-ssh`` ``>=6.0.0``
``apache-airflow-providers-common-compat`` ``>=1.12.0``
diff --git a/providers/sftp/pyproject.toml b/providers/sftp/pyproject.toml
index 469b7c0feff..c1cbbcf0905 100644
--- a/providers/sftp/pyproject.toml
+++ b/providers/sftp/pyproject.toml
@@ -59,7 +59,6 @@ requires-python = ">=3.10"
# Make sure to run ``prek update-providers-dependencies --all-files``
# After you modify the dependencies, and rebuild your Breeze CI image with
``breeze ci-image build``
dependencies = [
- "aiofiles>=23.2.0",
"apache-airflow>=2.11.0",
"apache-airflow-providers-ssh>=6.0.0",
"apache-airflow-providers-common-compat>=1.12.0",
diff --git a/providers/sftp/src/airflow/providers/sftp/exceptions.py
b/providers/sftp/src/airflow/providers/sftp/exceptions.py
index f256e6b7634..7f21b6e7b41 100644
--- a/providers/sftp/src/airflow/providers/sftp/exceptions.py
+++ b/providers/sftp/src/airflow/providers/sftp/exceptions.py
@@ -21,3 +21,7 @@ from airflow.providers.common.compat.sdk import
AirflowException
class ConnectionNotOpenedException(AirflowException):
"""Thrown when a connection has not been opened and has been tried to be
used."""
+
+
+class SFTPOperationError(AirflowException):
+ """Thrown when an SFTP operation (GET, PUT, DELETE) fails during
execution."""
diff --git a/providers/sftp/src/airflow/providers/sftp/hooks/sftp.py
b/providers/sftp/src/airflow/providers/sftp/hooks/sftp.py
index 4044b39bb0f..f89ff3553a8 100644
--- a/providers/sftp/src/airflow/providers/sftp/hooks/sftp.py
+++ b/providers/sftp/src/airflow/providers/sftp/hooks/sftp.py
@@ -19,27 +19,29 @@
from __future__ import annotations
+import asyncio
import concurrent.futures
import datetime
import functools
+import inspect
import os
import posixpath
import stat
import warnings
-from collections.abc import Callable, Generator, Sequence
-from contextlib import contextmanager, suppress
+from collections.abc import AsyncGenerator, Callable, Generator, Sequence
+from contextlib import AsyncExitStack, asynccontextmanager, contextmanager,
suppress
+from enum import Enum
from fnmatch import fnmatch
from io import BytesIO
from pathlib import Path, PurePosixPath
from typing import IO, TYPE_CHECKING, Any, cast
-import aiofiles
import asyncssh
from paramiko.config import SSH_PORT
from airflow.exceptions import AirflowProviderDeprecationWarning
from airflow.providers.common.compat.connection import get_async_connection
-from airflow.providers.common.compat.sdk import AirflowException, BaseHook,
Connection
+from airflow.providers.common.compat.sdk import AirflowException, BaseHook,
Connection, timezone
from airflow.providers.sftp.exceptions import ConnectionNotOpenedException
from airflow.providers.ssh.hooks.ssh import SSHHook
@@ -51,7 +53,31 @@ if TYPE_CHECKING:
CHUNK_SIZE = 64 * 1024 # 64KB
+class SFTPOperation(str, Enum):
+ """SFTP operation constants."""
+
+ GET = "get"
+ PUT = "put"
+ DELETE = "delete"
+
+
def handle_connection_management(func: Callable) -> Callable:
+ """
+ Run the wrapped hook method inside the hook's managed connection.
+
+ Both :class:`SFTPHook` and :class:`SFTPHookAsync` expose
``get_managed_conn()``, which
+ opens the connection on first entry and reuses it for nested entries, so a
decorated
+ method calling other decorated methods shares one connection with them.
+ """
+ if inspect.iscoroutinefunction(func):
+
+ @functools.wraps(func)
+ async def handle_async_connection_management_wrapper(self, *args: Any,
**kwargs: Any) -> Any:
+ async with self.get_managed_conn():
+ return await func(self, *args, **kwargs)
+
+ return handle_async_connection_management_wrapper
+
@functools.wraps(func)
def handle_connection_management_wrapper(self, *args: Any, **kwargs:
dict[str, Any]) -> Any:
if not self.use_managed_conn:
@@ -384,25 +410,6 @@ class SFTPHook(SSHHook):
"""
self.conn.remove(path) # type: ignore[arg-type, union-attr]
- @staticmethod
- def _validate_within_directory(base_dir: str, candidate: str) -> str:
- """
- Ensure ``candidate`` resolves to a path inside ``base_dir``.
-
- Directory-entry names are returned by the remote SFTP server and may
- contain ``..`` components; joining them into the local destination path
- could otherwise write outside it. Containment is verified before any
- local write or ``mkdir``.
- """
- base_real = os.path.realpath(base_dir)
- candidate_real = os.path.realpath(candidate)
- if candidate_real != base_real and os.path.commonpath([base_real,
candidate_real]) != base_real:
- raise ValueError(
- f"Refusing to write outside the destination directory: "
- f"{candidate!r} resolves outside {base_dir!r}"
- )
- return candidate
-
def retrieve_directory(self, remote_full_path: str, local_full_path: str,
prefetch: bool = True) -> None:
"""
Transfer the remote directory to a local location.
@@ -416,17 +423,16 @@ class SFTPHook(SSHHook):
"""
if Path(local_full_path).exists():
raise AirflowException(f"{local_full_path} already exists")
- Path(local_full_path).mkdir(parents=True)
+ dest = Path(local_full_path).resolve()
+ dest.mkdir(parents=True)
files, dirs, _ = self.get_tree_map(remote_full_path)
for dir_path in dirs:
- new_local_path = self._validate_within_directory(
- local_full_path, os.path.join(local_full_path,
os.path.relpath(dir_path, remote_full_path))
- )
+ new_local_path = str(dest / os.path.relpath(dir_path,
remote_full_path))
+ self._validate_within_directory(str(dest), new_local_path)
Path(new_local_path).mkdir(parents=True, exist_ok=True)
for file_path in files:
- new_local_path = self._validate_within_directory(
- local_full_path, os.path.join(local_full_path,
os.path.relpath(file_path, remote_full_path))
- )
+ new_local_path = str(dest / os.path.relpath(file_path,
remote_full_path))
+ self._validate_within_directory(str(dest), new_local_path)
self.retrieve_file(file_path, new_local_path, prefetch)
def retrieve_directory_concurrently(
@@ -461,19 +467,14 @@ class SFTPHook(SSHHook):
new_local_file_paths, remote_file_paths = [], []
files, dirs, _ = self.get_tree_map(remote_full_path)
for dir_path in dirs:
- new_local_path = self._validate_within_directory(
- local_full_path,
- os.path.join(local_full_path, os.path.relpath(dir_path,
remote_full_path)),
- )
+ new_local_path = os.path.join(local_full_path,
os.path.relpath(dir_path, remote_full_path))
+ self._validate_within_directory(local_full_path,
new_local_path)
Path(new_local_path).mkdir(parents=True, exist_ok=True)
for file in files:
+ new_local_path = os.path.join(local_full_path,
os.path.relpath(file, remote_full_path))
+ self._validate_within_directory(local_full_path,
new_local_path)
remote_file_paths.append(file)
- new_local_file_paths.append(
- self._validate_within_directory(
- local_full_path,
- os.path.join(local_full_path, os.path.relpath(file,
remote_full_path)),
- )
- )
+ new_local_file_paths.append(new_local_path)
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)
@@ -625,6 +626,27 @@ class SFTPHook(SSHHook):
return False
return True
+ @staticmethod
+ def _validate_within_directory(base: str, target: str) -> str:
+ """
+ Validate that target path is within the base directory.
+
+ Prevents directory traversal attacks.
+
+ :param base: The base/destination directory path
+ :param target: The target path to validate
+ :return: The target path if valid
+ :raises ValueError: If target path escapes the base directory
+ """
+ base_real = os.path.realpath(os.path.expanduser(base))
+ target_real = os.path.realpath(os.path.expanduser(target))
+
+ # Ensure target is within base directory
+ if not (target_real == base_real or target_real.startswith(base_real +
os.sep)):
+ raise ValueError(f"Path {target} is outside the destination
directory {base}")
+
+ return target
+
def walktree(
self,
path: str,
@@ -735,6 +757,74 @@ class SFTPHook(SSHHook):
return matched_files
+ def transfer(
+ self,
+ operation: str,
+ local_filepath: str | list[str] | None,
+ remote_filepath: str | list[str],
+ confirm: bool = True,
+ create_intermediate_dirs: bool = False,
+ concurrency: int = 1,
+ prefetch: bool = True,
+ ) -> None:
+ """
+ Perform a synchronous SFTP transfer operation (GET, PUT, or DELETE).
+
+ Centralizes transfer logic so both the operator and the trigger
+ can delegate to the hook, in line with the DRY principle.
+
+ :param operation: The SFTP operation - put, get, or delete.
+ :param local_filepath: Local file path(s).
+ :param remote_filepath: Remote file path(s).
+ :param confirm: Whether to confirm file size after PUT (default: True).
+ :param create_intermediate_dirs: Create missing intermediate
directories (default: False).
+ :param concurrency: Number of threads for directory transfers
(default: 1).
+ :param prefetch: Whether to prefetch during GET (default: True).
+ """
+ if isinstance(local_filepath, str):
+ local_filepath_array = [local_filepath] if local_filepath else []
+ else:
+ local_filepath_array = local_filepath or []
+
+ if isinstance(remote_filepath, str):
+ remote_filepath_array = [remote_filepath]
+ else:
+ remote_filepath_array = list(remote_filepath)
+
+ if operation.lower() == SFTPOperation.GET:
+ for local, remote in zip(local_filepath_array,
remote_filepath_array):
+ if create_intermediate_dirs:
+ Path(os.path.dirname(local)).mkdir(parents=True,
exist_ok=True)
+ if self.isdir(remote):
+ if concurrency > 1:
+ self.retrieve_directory_concurrently(
+ remote, local, workers=concurrency,
prefetch=prefetch
+ )
+ else:
+ self.retrieve_directory(remote, local,
prefetch=prefetch)
+ else:
+ self.retrieve_file(remote, local, prefetch=prefetch)
+ elif operation.lower() == SFTPOperation.PUT:
+ for local, remote in zip(local_filepath_array,
remote_filepath_array):
+ if create_intermediate_dirs:
+ self.create_directory(os.path.dirname(remote))
+ if os.path.isdir(local):
+ if concurrency > 1:
+ self.store_directory_concurrently(remote, local,
confirm=confirm, workers=concurrency)
+ else:
+ self.store_directory(remote, local, confirm=confirm)
+ else:
+ self.store_file(remote, local, confirm=confirm)
+ elif operation.lower() == SFTPOperation.DELETE:
+ for remote in remote_filepath_array:
+ if self.isdir(remote):
+ self.delete_directory(remote, include_files=True)
+ else:
+ try:
+ self.delete_file(remote)
+ except FileNotFoundError:
+ self.log.warning("Remote file %s does not exist.
Skipping delete.", remote)
+
class SFTPHookAsync(BaseHook):
"""
@@ -778,6 +868,10 @@ class SFTPHookAsync(BaseHook):
self.key_file = key_file
self.passphrase = passphrase
self.private_key = private_key
+ self.conn: asyncssh.SFTPClient | None = None
+ self._conn_count = 0
+ self._conn_lock = asyncio.Lock()
+ self._conn_stack: AsyncExitStack | None = None
def _parse_extras(self, conn: Connection) -> None:
"""Parse extra fields from the connection into instance fields."""
@@ -859,11 +953,55 @@ class SFTPHookAsync(BaseHook):
ssh_client_conn = await asyncssh.connect(**conn_config)
return ssh_client_conn
+ @asynccontextmanager
+ async def get_managed_conn(self) -> AsyncGenerator[asyncssh.SFTPClient]:
+ """
+ Context manager sharing one SSH connection and SFTP client across
nested uses.
+
+ The connection is opened on the first entry, reused by any entry made
while it is
+ still open, and closed when the last user exits, mirroring
:meth:`SFTPHook.get_managed_conn`.
+ Hook methods are wrapped in it, so wrapping several calls in this
context manager makes
+ them run over a single connection.
+ """
+ async with self._conn_lock:
+ if self.conn is None:
+ stack = AsyncExitStack()
+ try:
+ ssh_conn = await stack.enter_async_context(await
self._get_conn())
+ self.conn = await
stack.enter_async_context(ssh_conn.start_sftp_client())
+ except BaseException:
+ await stack.aclose()
+ raise
+ self._conn_stack = stack
+ self._conn_count += 1
+ sftp = self.conn
+ try:
+ yield sftp
+ finally:
+ self._conn_count -= 1
+ if self._conn_count == 0:
+ open_stack, self._conn_stack, self.conn = self._conn_stack,
None, None
+ if open_stack is not None:
+ await open_stack.aclose()
+
+ def get_conn_count(self) -> int:
+ """Get the number of users currently sharing the open connection."""
+ return self._conn_count
+
+ def _get_open_conn(self) -> asyncssh.SFTPClient:
+ if self.conn is None:
+ raise ConnectionNotOpenedException(
+ "Connection not open, use `async with hook.get_managed_conn()`
to open it first."
+ )
+ return self.conn
+
+ @handle_connection_management
async def retrieve_file(
self,
remote_full_path: str,
local_full_path: str | os.PathLike[str] | IO[bytes],
chunk_size: int = CHUNK_SIZE,
+ prefetch: bool = True,
) -> None:
"""
Transfer the remote file to a local location asynchronously.
@@ -874,28 +1012,33 @@ class SFTPHookAsync(BaseHook):
:param remote_full_path: Full path to the remote file.
:param local_full_path: Full path to the local file or a binary
file-like buffer.
:param chunk_size: Size of chunks to read at a time (default: 64KB).
+ :param prefetch: Whether to allow read-ahead requests to be sent
concurrently (default: True). When
+ ``False``, only one request is kept in flight at a time, mirroring
+ :meth:`SFTPHook.retrieve_file`'s ``prefetch`` semantics.
"""
- async with await self._get_conn() as ssh_conn:
- async with ssh_conn.start_sftp_client() as sftp:
- async with sftp.open(remote_full_path, "rb") as remote_file:
- if isinstance(local_full_path, (str, os.PathLike)):
- async with aiofiles.open(local_full_path, "wb") as f:
- while True:
- chunk = await remote_file.read(chunk_size)
- if not chunk:
- break
- await f.write(cast("bytes", chunk))
- else:
- while True:
- chunk = await remote_file.read(chunk_size)
- if not chunk:
- break
- local_full_path.write(cast("bytes", chunk))
- if hasattr(local_full_path, "seek"):
- local_full_path.seek(0)
+ sftp = self._get_open_conn()
+ if isinstance(local_full_path, (str, os.PathLike)):
+ get_kwargs: dict[str, Any] = {"block_size": chunk_size}
+ if not prefetch:
+ get_kwargs["max_requests"] = 1
+ await sftp.get(remote_full_path, os.fspath(local_full_path),
**get_kwargs)
+ return
+
+ async with sftp.open(remote_full_path, "rb") as remote_file:
+ while True:
+ chunk = await remote_file.read(chunk_size)
+ if not chunk:
+ break
+ local_full_path.write(cast("bytes", chunk))
+ if hasattr(local_full_path, "seek"):
+ local_full_path.seek(0)
+ @handle_connection_management
async def store_file(
- self, remote_full_path: str, local_full_path: str | os.PathLike[str] |
IO[bytes]
+ self,
+ remote_full_path: str,
+ local_full_path: str | os.PathLike[str] | IO[bytes],
+ confirm: bool = True,
) -> None:
"""
Transfer a local file to the remote location.
@@ -908,31 +1051,40 @@ class SFTPHookAsync(BaseHook):
:param remote_full_path: full path to the remote file
:param local_full_path: full path to the local file or a binary
file-like buffer
+ :param confirm: whether to verify the remote file size matches the
local size after
+ upload (default: True), mirroring :meth:`SFTPHook.store_file`'s
``confirm`` semantics.
"""
if isinstance(local_full_path, bytes):
raise TypeError("Unsupported type for local_full_path: bytes. Wrap
raw bytes in BytesIO.")
- async with await self._get_conn() as ssh_conn:
- async with ssh_conn.start_sftp_client() as sftp:
- with suppress(asyncssh.SFTPFailure):
- remote_path = PurePosixPath(remote_full_path)
- await sftp.makedirs(str(remote_path.parent))
-
- if isinstance(local_full_path, (str, os.PathLike)):
- await sftp.put(str(local_full_path), remote_full_path)
- elif hasattr(local_full_path, "read"):
- async with sftp.open(remote_full_path, "wb") as f:
- stream = local_full_path
- if hasattr(stream, "seek"):
- stream.seek(0)
- data = stream.read()
- await f.write(data)
- else:
- raise TypeError(
- f"Unsupported type for local_full_path:
{type(local_full_path)}. "
- "Expected a binary file-like object or a path-like
object."
- )
+ sftp = self._get_open_conn()
+ with suppress(asyncssh.SFTPFailure):
+ remote_path = PurePosixPath(remote_full_path)
+ await sftp.makedirs(str(remote_path.parent))
+
+ if isinstance(local_full_path, (str, os.PathLike)):
+ await sftp.put(str(local_full_path), remote_full_path)
+ uploaded_size = await asyncio.to_thread(os.path.getsize,
local_full_path)
+ elif hasattr(local_full_path, "read"):
+ async with sftp.open(remote_full_path, "wb") as f:
+ stream = local_full_path
+ if hasattr(stream, "seek"):
+ stream.seek(0)
+ data = stream.read()
+ await f.write(data)
+ uploaded_size = len(data)
+ else:
+ raise TypeError(
+ f"Unsupported type for local_full_path:
{type(local_full_path)}. "
+ "Expected a binary file-like object or a path-like object."
+ )
+ if confirm:
+ remote_attrs = await sftp.stat(remote_full_path)
+ if remote_attrs.size != uploaded_size:
+ raise OSError(f"size mismatch in put! {remote_attrs.size} !=
{uploaded_size}")
+
+ @handle_connection_management
async def mkdir(self, path: str) -> None:
"""
Create a directory on the remote system asynchronously.
@@ -941,10 +1093,9 @@ class SFTPHookAsync(BaseHook):
:param path: Full path to the remote directory to create.
"""
- async with await self._get_conn() as ssh_conn:
- async with ssh_conn.start_sftp_client() as sftp:
- await sftp.makedirs(path)
+ await self._get_open_conn().makedirs(path)
+ @handle_connection_management
async def list_directory(self, path: str = "", recursive: bool = False) ->
list[str] | None:
"""
List files in a directory on the remote system asynchronously.
@@ -974,16 +1125,13 @@ class SFTPHookAsync(BaseHook):
return None
return sorted(files)
- async with await self._get_conn() as ssh_conn:
- async with ssh_conn.start_sftp_client() as sftp:
- try:
- entries = await sftp.readdir(path)
- except asyncssh.SFTPNoSuchFile:
- return None
- return sorted(os.fsdecode(entry.filename) for entry in entries)
-
- return None
+ try:
+ entries = await self._get_open_conn().readdir(path)
+ except asyncssh.SFTPNoSuchFile:
+ return None
+ return sorted(os.fsdecode(entry.filename) for entry in entries)
+ @handle_connection_management
async def walktree(
self,
path: str,
@@ -998,54 +1146,52 @@ class SFTPHookAsync(BaseHook):
This mirrors :meth:`SFTPHook.walktree` contract and calls callback
functions for
regular files, directories, and unknown file types.
"""
- async with await self._get_conn() as ssh_conn:
- async with ssh_conn.start_sftp_client() as sftp:
- visited_dirs: set[str] = set()
-
- async def _canonical_dir(dir_path: str) -> str:
- with suppress(asyncssh.SFTPError):
- return os.fsdecode(await sftp.realpath(dir_path))
- return posixpath.normpath(dir_path)
-
- async def _walk(dir_path: str) -> None:
- canonical_dir = await _canonical_dir(dir_path)
- if canonical_dir in visited_dirs:
- return
- visited_dirs.add(canonical_dir)
+ sftp = self._get_open_conn()
+ visited_dirs: set[str] = set()
- try:
- entries = await sftp.readdir(dir_path)
- except asyncssh.SFTPNoSuchFile:
- # Directory may disappear mid-walk on busy drops; skip
and continue.
- return
+ async def _canonical_dir(dir_path: str) -> str:
+ with suppress(asyncssh.SFTPError):
+ return os.fsdecode(await sftp.realpath(dir_path))
+ return posixpath.normpath(dir_path)
- for entry in sorted(entries, key=lambda file:
os.fsdecode(file.filename)):
- filename = os.fsdecode(entry.filename)
- if filename in {".", ".."}:
- continue
+ async def _walk(dir_path: str) -> None:
+ canonical_dir = await _canonical_dir(dir_path)
+ if canonical_dir in visited_dirs:
+ return
+ visited_dirs.add(canonical_dir)
- pathname = posixpath.join(dir_path, filename)
- permissions = entry.attrs.permissions
-
- if permissions is not None and
stat.S_ISDIR(permissions):
- dcallback(pathname)
- if recurse:
- await _walk(pathname)
- elif permissions is not None and
stat.S_ISREG(permissions):
- fcallback(pathname)
- else:
- ucallback(pathname)
+ try:
+ entries = await sftp.readdir(dir_path)
+ except asyncssh.SFTPNoSuchFile:
+ # Directory may disappear mid-walk on busy drops; skip and
continue.
+ return
+
+ for entry in sorted(entries, key=lambda file:
os.fsdecode(file.filename)):
+ filename = os.fsdecode(entry.filename)
+ if filename in {".", ".."}:
+ continue
+
+ pathname = posixpath.join(dir_path, filename)
+ permissions = entry.attrs.permissions
+
+ if permissions is not None and stat.S_ISDIR(permissions):
+ dcallback(pathname)
+ if recurse:
+ await _walk(pathname)
+ elif permissions is not None and stat.S_ISREG(permissions):
+ fcallback(pathname)
+ else:
+ ucallback(pathname)
- await _walk(path)
+ await _walk(path)
- async def read_directory(self, path: str = "") ->
Sequence[asyncssh.sftp.SFTPName] | None: # type: ignore[return]
+ @handle_connection_management
+ async def read_directory(self, path: str = "") ->
Sequence[asyncssh.sftp.SFTPName] | None:
"""Return a list of files along with their attributes on the SFTP
server at the provided path."""
- async with await self._get_conn() as ssh_conn:
- async with ssh_conn.start_sftp_client() as sftp:
- try:
- return await sftp.readdir(path)
- except asyncssh.SFTPNoSuchFile:
- return None
+ try:
+ return await self._get_open_conn().readdir(path)
+ except asyncssh.SFTPNoSuchFile:
+ return None
async def get_files_and_attrs_by_pattern(
self, path: str = "", fnmatch_pattern: str = ""
@@ -1061,7 +1207,8 @@ class SFTPHookAsync(BaseHook):
matched_files = [file for file in files_list if
fnmatch(str(file.filename), fnmatch_pattern)]
return matched_files
- async def get_mod_time(self, path: str) -> str: # type: ignore[return]
+ @handle_connection_management
+ async def get_mod_time(self, path: str) -> str:
"""
Make SFTP async connection.
@@ -1070,13 +1217,300 @@ class SFTPHookAsync(BaseHook):
:param path: full path to the remote file
"""
- async with await self._get_conn() as ssh_conn:
- async with ssh_conn.start_sftp_client() as sftp:
- try:
- ftp_mdtm = await sftp.stat(path)
- modified_time = ftp_mdtm.mtime
- mod_time =
datetime.datetime.fromtimestamp(modified_time).strftime("%Y%m%d%H%M%S") #
type: ignore[arg-type]
- self.log.info("Found File %s last modified: %s",
str(path), str(mod_time))
- return mod_time
- except asyncssh.SFTPNoSuchFile:
- raise AirflowException("No files matching")
+ try:
+ ftp_mdtm = await self._get_open_conn().stat(path)
+ except asyncssh.SFTPNoSuchFile:
+ raise AirflowException("No files matching")
+ modified_time = ftp_mdtm.mtime
+ mod_time =
datetime.datetime.fromtimestamp(modified_time).strftime("%Y%m%d%H%M%S") #
type: ignore[arg-type]
+ self.log.info("Found File %s last modified: %s", str(path),
str(mod_time))
+ return mod_time
+
+ async def sense_files_by_pattern(
+ self,
+ path: str,
+ fnmatch_pattern: str,
+ newer_than: datetime.datetime | None = None,
+ ) -> list[str]:
+ """
+ Return the names of files at ``path`` matching ``fnmatch_pattern``.
+
+ If ``newer_than`` is provided, only files modified after that
timestamp are returned; files
+ without a reported modification time are skipped in that case.
+
+ :param path: directory on the SFTP server to search for files matching
the pattern
+ :param fnmatch_pattern: pattern used to match filenames, see the
``fnmatch`` std library module
+ :param newer_than: if provided, only files modified after this UTC
timestamp are returned
+ """
+ files = await self.get_files_and_attrs_by_pattern(path=path,
fnmatch_pattern=fnmatch_pattern)
+ if not newer_than:
+ return [str(file.filename) for file in files]
+
+ matched_files = []
+ for file in files:
+ if file.attrs.mtime is None:
+ continue
+ if newer_than <= self._mod_time_to_utc(file.attrs.mtime):
+ matched_files.append(str(file.filename))
+ return matched_files
+
+ async def sense_path(self, path: str, newer_than: datetime.datetime | None
= None) -> bool:
+ """
+ Return whether ``path`` exists and, if ``newer_than`` is provided, was
modified since.
+
+ :param path: full path to the remote file
+ :param newer_than: if provided, the file must have been modified after
this UTC timestamp
+ """
+ mod_time = await self.get_mod_time(path)
+ if not newer_than:
+ return True
+ return newer_than <= self._mod_time_to_utc(mod_time)
+
+ @staticmethod
+ def _mod_time_to_utc(mod_time: int | float | str) -> datetime.datetime:
+ """Convert a modification time, either an epoch timestamp or
``%Y%m%d%H%M%S`` string, to UTC."""
+ if not isinstance(mod_time, str):
+ mod_time =
datetime.datetime.fromtimestamp(float(mod_time)).strftime("%Y%m%d%H%M%S")
+ return timezone.convert_to_utc(datetime.datetime.strptime(mod_time,
"%Y%m%d%H%M%S"))
+
+ @handle_connection_management
+ async def isdir(self, path: str) -> bool:
+ """
+ Check if the path provided is a directory.
+
+ :param path: full path to the remote directory to check
+ """
+ try:
+ attrs = await self._get_open_conn().stat(path)
+ except asyncssh.SFTPNoSuchFile:
+ return False
+ return attrs.permissions is not None and
stat.S_ISDIR(attrs.permissions)
+
+ @handle_connection_management
+ async def path_exists(self, path: str) -> bool:
+ """
+ Whether a remote entity exists.
+
+ :param path: full path to the remote file or directory
+ """
+ try:
+ await self._get_open_conn().stat(path)
+ except asyncssh.SFTPNoSuchFile:
+ return False
+ return True
+
+ @handle_connection_management
+ async def create_directory(self, path: str) -> None:
+ """
+ Create a directory (and any missing parents) on the remote system
asynchronously.
+
+ Returns silently if the target directory already exists, mirroring
+ :meth:`SFTPHook.create_directory`.
+
+ :param path: full path to the remote directory to create
+ """
+ await self._get_open_conn().makedirs(path, exist_ok=True)
+
+ @handle_connection_management
+ async def delete_file(self, path: str) -> None:
+ """
+ Remove a file on the server asynchronously.
+
+ :param path: full path to the remote file
+ """
+ await self._get_open_conn().unlink(path)
+
+ @handle_connection_management
+ async def delete_directory(self, path: str, include_files: bool = False)
-> None:
+ """
+ Delete a directory on the remote system asynchronously.
+
+ :param path: full path to the remote directory to delete
+ :param include_files: whether to recursively delete the directory's
contents first
+ """
+ files: list[str] = []
+ dirs: list[str] = []
+
+ if include_files:
+ files, dirs, _ = await self.get_tree_map(path)
+ dirs = dirs[::-1] # reverse the order for deleting deepest
directories first
+
+ sftp = self._get_open_conn()
+ for file_path in files:
+ await sftp.remove(file_path)
+ for dir_path in dirs:
+ await sftp.rmdir(dir_path)
+ await sftp.rmdir(path)
+
+ @handle_connection_management
+ async def get_tree_map(
+ self, path: str, prefix: str | None = None, delimiter: str | None =
None
+ ) -> tuple[list[str], list[str], list[str]]:
+ """
+ Get tuple with recursive lists of files, directories and unknown paths
asynchronously.
+
+ It is possible to filter results by giving prefix and/or delimiter
parameters.
+
+ :param path: path from which tree will be built
+ :param prefix: if set paths will be added if start with prefix
+ :param delimiter: if set paths will be added if end with delimiter
+ :return: tuple with list of files, dirs and unknown items
+ """
+ files: list[str] = []
+ dirs: list[str] = []
+ unknowns: list[str] = []
+
+ def append_matching_path_callback(list_: list[str]) -> Callable:
+ return lambda item: (
+ list_.append(item) if SFTPHook._is_path_match(item, prefix,
delimiter) else None
+ )
+
+ await self.walktree(
+ path=path,
+ fcallback=append_matching_path_callback(files),
+ dcallback=append_matching_path_callback(dirs),
+ ucallback=append_matching_path_callback(unknowns),
+ recurse=True,
+ )
+
+ return files, dirs, unknowns
+
+ @handle_connection_management
+ async def retrieve_directory(
+ self, remote_full_path: str, local_full_path: str, prefetch: bool =
True
+ ) -> None:
+ """
+ Transfer the remote directory to a local location asynchronously.
+
+ The whole tree is walked and downloaded over a single connection.
+
+ :param remote_full_path: full path to the remote directory
+ :param local_full_path: full path to the local directory
+ :param prefetch: whether read-ahead requests are sent concurrently
(default: True)
+ """
+ if await asyncio.to_thread(Path(local_full_path).exists):
+ raise FileExistsError(f"{local_full_path} already exists")
+ dest = await asyncio.to_thread(Path(local_full_path).resolve)
+ await asyncio.to_thread(dest.mkdir, parents=True)
+ files, dirs, _ = await self.get_tree_map(remote_full_path)
+ remote_base = PurePosixPath(remote_full_path)
+ for dir_path in dirs:
+ relative_path = PurePosixPath(dir_path).relative_to(remote_base)
+ new_local_path = str(dest / relative_path)
+ SFTPHook._validate_within_directory(str(dest), new_local_path)
+ await asyncio.to_thread(Path(new_local_path).mkdir, parents=True,
exist_ok=True)
+ for file_path in files:
+ relative_path = PurePosixPath(file_path).relative_to(remote_base)
+ new_local_path = str(dest / relative_path)
+ SFTPHook._validate_within_directory(str(dest), new_local_path)
+ await self.retrieve_file(file_path, new_local_path,
prefetch=prefetch)
+
+ @handle_connection_management
+ async def store_directory(
+ self, remote_full_path: str, local_full_path: str, confirm: bool = True
+ ) -> None:
+ """
+ Transfer a local directory to the remote location asynchronously.
+
+ The whole tree is created and uploaded over a single connection.
+
+ :param remote_full_path: full path to the remote directory
+ :param local_full_path: full path to the local directory
+ :param confirm: whether to verify each uploaded file's size (default:
True)
+ """
+ if await self.path_exists(remote_full_path):
+ raise FileExistsError(f"{remote_full_path} already exists")
+ await self.create_directory(remote_full_path)
+ entries = await asyncio.to_thread(lambda:
list(os.walk(local_full_path)))
+ local_base = Path(local_full_path)
+ for root, dirs, files in entries:
+ for dir_name in dirs:
+ dir_path = Path(root) / dir_name
+ relative_path = dir_path.relative_to(local_base).as_posix()
+ await
self.create_directory(str(PurePosixPath(remote_full_path) / relative_path))
+ for file_name in files:
+ file_path = Path(root) / file_name
+ relative_path = file_path.relative_to(local_base).as_posix()
+ new_remote_path = str(PurePosixPath(remote_full_path) /
relative_path)
+ await self.store_file(new_remote_path, str(file_path),
confirm=confirm)
+
+ @handle_connection_management
+ async def transfer(
+ self,
+ operation: str,
+ local_filepath: str | list[str] | None,
+ remote_filepath: str | list[str],
+ confirm: bool = True,
+ create_intermediate_dirs: bool = False,
+ concurrency: int = 1,
+ prefetch: bool = True,
+ ) -> None:
+ """
+ Perform an SFTP transfer operation (GET, PUT, or DELETE) using native
async I/O.
+
+ Mirrors :meth:`SFTPHook.transfer`, including directory transfers,
missing-file
+ deletes, and the ``confirm``/``prefetch`` options, so
``deferrable=True`` behaves
+ the same as the synchronous path. The whole call runs over a single
connection and
+ ``concurrency`` bounds how many top-level paths are in flight on it at
once; files
+ inside a directory are transferred sequentially. The synchronous path
differs there:
+ it opens one connection per worker and transfers directory contents
concurrently.
+ """
+ if isinstance(local_filepath, str):
+ local_filepath_array = [local_filepath] if local_filepath else []
+ else:
+ local_filepath_array = local_filepath or []
+
+ if isinstance(remote_filepath, str):
+ remote_filepath_array = [remote_filepath]
+ else:
+ remote_filepath_array = list(remote_filepath)
+
+ semaphore = asyncio.Semaphore(concurrency)
+
+ async def _bounded(coro):
+ async with semaphore:
+ return await coro
+
+ if operation.lower() == SFTPOperation.GET:
+
+ async def _get(local: str, remote: str):
+ if create_intermediate_dirs:
+ await
asyncio.to_thread(Path(os.path.dirname(local)).mkdir, parents=True,
exist_ok=True)
+ if await self.isdir(remote):
+ await self.retrieve_directory(remote, local,
prefetch=prefetch)
+ else:
+ await self.retrieve_file(remote, local, prefetch=prefetch)
+
+ tasks = [
+ asyncio.create_task(_bounded(_get(local, remote)))
+ for local, remote in zip(local_filepath_array,
remote_filepath_array)
+ ]
+ await asyncio.gather(*tasks)
+ elif operation.lower() == SFTPOperation.PUT:
+
+ async def _put(local: str, remote: str):
+ if create_intermediate_dirs:
+ await self.create_directory(os.path.dirname(remote))
+ if await asyncio.to_thread(os.path.isdir, local):
+ await self.store_directory(remote, local, confirm=confirm)
+ else:
+ await self.store_file(remote, local, confirm=confirm)
+
+ tasks = [
+ asyncio.create_task(_bounded(_put(local, remote)))
+ for local, remote in zip(local_filepath_array,
remote_filepath_array)
+ ]
+ await asyncio.gather(*tasks)
+ elif operation.lower() == SFTPOperation.DELETE:
+
+ async def _delete(remote: str):
+ if await self.isdir(remote):
+ await self.delete_directory(remote, include_files=True)
+ else:
+ try:
+ await self.delete_file(remote)
+ except asyncssh.SFTPNoSuchFile:
+ self.log.warning("Remote file %s does not exist.
Skipping delete.", remote)
+
+ tasks = [asyncio.create_task(_bounded(_delete(remote))) for remote
in remote_filepath_array]
+ await asyncio.gather(*tasks)
diff --git a/providers/sftp/src/airflow/providers/sftp/operators/sftp.py
b/providers/sftp/src/airflow/providers/sftp/operators/sftp.py
index 0b47b5b7d5d..b34ae6ad7ff 100644
--- a/providers/sftp/src/airflow/providers/sftp/operators/sftp.py
+++ b/providers/sftp/src/airflow/providers/sftp/operators/sftp.py
@@ -19,44 +19,37 @@
from __future__ import annotations
-import errno
-import os
import socket
from collections.abc import Sequence
-from pathlib import Path
from typing import Any
import paramiko
-from airflow.providers.common.compat.sdk import AirflowException, BaseOperator
-from airflow.providers.sftp.hooks.sftp import SFTPHook
-
-
-class SFTPOperation:
- """Operation that can be used with SFTP."""
-
- PUT = "put"
- GET = "get"
- DELETE = "delete"
+from airflow.providers.common.compat.sdk import AirflowException,
BaseOperator, conf
+from airflow.providers.sftp.exceptions import SFTPOperationError
+from airflow.providers.sftp.hooks.sftp import SFTPHook, SFTPOperation
class SFTPOperator(BaseOperator):
"""
SFTPOperator for transferring files from remote host to local or vice a
versa.
- This operator uses sftp_hook to open sftp transport channel that serve as
basis for file transfer.
+ This operator uses sftp_hook to open an SFTP transport channel that serves
as
+ the basis for file transfer. All transfer logic is delegated to
+ ``SFTPHook.transfer()`` so that both the synchronous and deferrable code
paths
+ share a single, authoritative implementation.
:param ssh_conn_id: :ref:`ssh connection id<howto/connection:ssh>`
from airflow Connections.
- :param sftp_hook: predefined SFTPHook to use
+ :param sftp_hook: predefined SFTPHook to use.
Either `sftp_hook` or `ssh_conn_id` needs to be provided.
- :param remote_host: remote host to connect (templated)
+ :param remote_host: remote host to connect (templated).
Nullable. If provided, it will replace the `remote_host` which was
defined in `sftp_hook` or predefined in the connection of
`ssh_conn_id`.
:param local_filepath: local file path or list of local file paths to get
or put. (templated)
:param remote_filepath: remote file path or list of remote file paths to
get, put, or delete. (templated)
- :param operation: specify operation 'get', 'put', or 'delete', defaults to
put
- :param confirm: specify if the SFTP operation should be confirmed,
defaults to True
+ :param operation: specify operation ``'get'``, ``'put'``, or ``'delete'``.
Defaults to ``'put'``.
+ :param confirm: specify if the SFTP operation should be confirmed.
Defaults to True.
:param create_intermediate_dirs: create missing intermediate directories
when
copying from remote to local and vice-versa. Default is False.
@@ -74,10 +67,15 @@ class SFTPOperator(BaseOperator):
create_intermediate_dirs=True,
dag=dag,
)
- :param concurrency: Number of threads when transferring directories. Each
thread opens a new SFTP connection.
- This parameter is used only when transferring directories, not
individual files. (Default is 1)
- :param prefetch: controls whether prefetch is performed (default: True)
+ :param concurrency: number of threads when transferring directories. Each
thread opens
+ a new SFTP connection. Only applies to directory transfers. (Default:
1)
+ :param prefetch: controls whether prefetch is performed on GET transfers.
(Default: True)
+ :param deferrable: run the operator in deferrable mode. When True, the
worker slot is
+ freed during the transfer and reclaimed only when the transfer
completes.
+ Best suited for single large file transfers. For bulk directory
transfers involving
+ many files, consider using ``async PythonOperator`` with
``SFTPClientPool`` instead,
+ which provides true async multiplexing via a single event loop.
(Default: False)
"""
template_fields: Sequence[str] = ("local_filepath", "remote_filepath",
"remote_host")
@@ -95,6 +93,7 @@ class SFTPOperator(BaseOperator):
create_intermediate_dirs: bool = False,
concurrency: int = 1,
prefetch: bool = True,
+ deferrable: bool = conf.getboolean("operators", "default_deferrable",
fallback=False),
**kwargs,
) -> None:
super().__init__(**kwargs)
@@ -108,20 +107,24 @@ class SFTPOperator(BaseOperator):
self.remote_filepath = remote_filepath
self.concurrency = concurrency
self.prefetch = prefetch
+ self.deferrable = deferrable
def execute(self, context: Any) -> str | list[str] | None:
+ local_filepath_array: list[str] = []
if self.local_filepath is None:
local_filepath_array = []
elif isinstance(self.local_filepath, str):
local_filepath_array = [self.local_filepath]
else:
- local_filepath_array = self.local_filepath
+ local_filepath_array = list(self.local_filepath)
- if isinstance(self.remote_filepath, str):
- remote_filepath_array = [self.remote_filepath]
- else:
- remote_filepath_array = self.remote_filepath
+ remote_filepath_array: list[str] = (
+ [self.remote_filepath] if isinstance(self.remote_filepath, str)
else list(self.remote_filepath)
+ )
+ # ------------------------------------------------------------------ #
+ # Input validation
#
+ # ------------------------------------------------------------------ #
if self.operation.lower() in (SFTPOperation.GET, SFTPOperation.PUT)
and len(
local_filepath_array
) != len(remote_filepath_array):
@@ -136,106 +139,91 @@ class SFTPOperator(BaseOperator):
if self.operation.lower() not in (SFTPOperation.GET,
SFTPOperation.PUT, SFTPOperation.DELETE):
raise TypeError(
f"Unsupported operation value {self.operation}, "
- f"expected {SFTPOperation.GET} or {SFTPOperation.PUT} or
{SFTPOperation.DELETE}."
+ f"expected {SFTPOperation.GET!r}, {SFTPOperation.PUT!r}, "
+ f"or {SFTPOperation.DELETE!r}."
)
if self.concurrency < 1:
- raise ValueError(f"concurrency should be greater than 0, got
{self.concurrency}")
+ raise ValueError(f"concurrency should be >= 1, got
{self.concurrency}")
- file_msg = None
- try:
- if self.remote_host is not None:
- self.log.info(
- "remote_host is provided explicitly. "
- "It will replace the remote_host which was defined "
- "in sftp_hook or predefined in connection of ssh_conn_id."
+ # ------------------------------------------------------------------ #
+ # Synchronous path — delegate all transfer logic to the hook #
+ # ------------------------------------------------------------------ #
+ if self.remote_host is not None:
+ self.log.info(
+ "remote_host is provided explicitly. "
+ "It will replace the remote_host which was defined "
+ "in sftp_hook or predefined in connection of ssh_conn_id."
+ )
+
+ if self.ssh_conn_id:
+ if self.sftp_hook and isinstance(self.sftp_hook, SFTPHook):
+ self.log.info("ssh_conn_id is ignored when sftp_hook is
provided.")
+ else:
+ self.log.info("sftp_hook not provided or invalid. Trying
ssh_conn_id to create SFTPHook.")
+ self.sftp_hook = SFTPHook(
+ ssh_conn_id=self.ssh_conn_id,
+ remote_host=self.remote_host or "",
)
- if self.ssh_conn_id:
- if self.sftp_hook and isinstance(self.sftp_hook, SFTPHook):
- self.log.info("ssh_conn_id is ignored when sftp_hook is
provided.")
- else:
- self.log.info("sftp_hook not provided or invalid. Trying
ssh_conn_id to create SFTPHook.")
- self.sftp_hook = SFTPHook(
- ssh_conn_id=self.ssh_conn_id,
remote_host=self.remote_host or ""
- )
-
- if not self.sftp_hook:
- raise AirflowException("Cannot operate without sftp_hook or
ssh_conn_id.")
-
- if self.operation.lower() in (SFTPOperation.GET,
SFTPOperation.PUT):
- for _local_filepath, _remote_filepath in
zip(local_filepath_array, remote_filepath_array):
- if self.operation.lower() == SFTPOperation.GET:
- local_folder = os.path.dirname(_local_filepath)
- if self.create_intermediate_dirs:
- Path(local_folder).mkdir(parents=True,
exist_ok=True)
- file_msg = f"from {_remote_filepath} to
{_local_filepath}"
- self.log.info("Starting to transfer %s", file_msg)
- if self.sftp_hook.isdir(_remote_filepath):
- if self.concurrency > 1:
- self.sftp_hook.retrieve_directory_concurrently(
- _remote_filepath,
- _local_filepath,
- workers=self.concurrency,
- prefetch=self.prefetch,
- )
- elif self.concurrency == 1:
-
self.sftp_hook.retrieve_directory(_remote_filepath, _local_filepath)
- else:
- self.sftp_hook.retrieve_file(_remote_filepath,
_local_filepath)
- elif self.operation.lower() == SFTPOperation.PUT:
- remote_folder = os.path.dirname(_remote_filepath)
- if self.create_intermediate_dirs:
- self.sftp_hook.create_directory(remote_folder)
- file_msg = f"from {_local_filepath} to
{_remote_filepath}"
- self.log.info("Starting to transfer file %s", file_msg)
- if os.path.isdir(_local_filepath):
- if self.concurrency > 1:
- self.sftp_hook.store_directory_concurrently(
- _remote_filepath,
- _local_filepath,
- confirm=self.confirm,
- workers=self.concurrency,
- )
- elif self.concurrency == 1:
- self.sftp_hook.store_directory(
- _remote_filepath, _local_filepath,
confirm=self.confirm
- )
- else:
- self.sftp_hook.store_file(_remote_filepath,
_local_filepath, confirm=self.confirm)
- elif self.operation.lower() == SFTPOperation.DELETE:
- for _remote_filepath in remote_filepath_array:
- file_msg = f"{_remote_filepath}"
- self.log.info("Starting to delete %s", file_msg)
- try:
- if self.sftp_hook.isdir(_remote_filepath):
- self.sftp_hook.delete_directory(_remote_filepath,
include_files=True)
- else:
- self.sftp_hook.delete_file(_remote_filepath)
- except OSError as exc:
- if self._is_missing_path_error(exc):
- self.log.warning(
- "Remote path %s does not exist. Skipping
delete.", _remote_filepath
- )
- continue
- raise
+ if not self.sftp_hook:
+ raise AirflowException("Cannot operate without sftp_hook or
ssh_conn_id.")
- except Exception as e:
- raise AirflowException(
- f"Error while processing {self.operation.upper()} operation
{file_msg}, error: {e}"
+ if self.deferrable:
+ from airflow.providers.sftp.triggers.sftp import
SFTPTransferTrigger
+
+ sftp_conn_id = self.ssh_conn_id or self.sftp_hook.ssh_conn_id
+ if not sftp_conn_id:
+ raise ValueError(
+ "deferrable=True requires a connection id: set ssh_conn_id
or pass an sftp_hook "
+ "that was created with one."
+ )
+ self.defer(
+ trigger=SFTPTransferTrigger(
+ sftp_conn_id=sftp_conn_id,
+ local_filepath=self.local_filepath,
+ remote_filepath=self.remote_filepath,
+ operation=self.operation,
+ confirm=self.confirm,
+ create_intermediate_dirs=self.create_intermediate_dirs,
+ remote_host=self.remote_host,
+ concurrency=self.concurrency,
+ prefetch=self.prefetch,
+ ),
+ method_name="execute_complete",
)
+ try:
+ for idx, remote_fp in enumerate(remote_filepath_array):
+ local_fp = local_filepath_array[idx] if local_filepath_array
else ""
+ self.sftp_hook.transfer(
+ local_filepath=local_fp,
+ remote_filepath=remote_fp,
+ operation=self.operation,
+ confirm=self.confirm,
+ create_intermediate_dirs=self.create_intermediate_dirs,
+ concurrency=self.concurrency,
+ prefetch=self.prefetch,
+ )
+ except Exception as e:
+ raise SFTPOperationError(
+ f"Error while processing {self.operation.upper()} operation,
error: {e}"
+ ) from e
+
return self.local_filepath
- @staticmethod
- def _is_missing_path_error(exc: Exception) -> bool:
- if isinstance(exc, FileNotFoundError):
- return True
- if isinstance(exc, OSError) and exc.errno == errno.ENOENT:
- return True
- if exc.args and isinstance(exc.args[0], int) and exc.args[0] ==
errno.ENOENT:
- return True
- return False
+ def execute_complete(self, context: Any, event: dict[str, Any]) -> str |
list[str] | None:
+ """
+ Handle completion from ``SFTPOperatorTrigger``.
+
+ :param context: Airflow task context
+ :param event: trigger result dict with ``status`` and ``message`` keys
+ :raises AirflowException: if the trigger reported an error
+ """
+ if event.get("status") == "error":
+ raise AirflowException(event.get("message", "Unknown error during
deferrable SFTP transfer"))
+ self.log.info("Deferrable SFTP transfer completed: %s",
event.get("message"))
+ return self.local_filepath
def get_openlineage_facets_on_start(self):
"""
@@ -279,10 +267,6 @@ class SFTPOperator(BaseOperator):
if hasattr(hook, "port"):
remote_port = hook.port
- # Since v4.1.0, SFTPOperator accepts both a string (single file) and a
list of
- # strings (multiple files) as local_filepath and remote_filepath, and
internally
- # keeps them as list in both cases. But before 4.1.0, only single
string is
- # allowed. So we consider both cases here for backward compatibility.
if isinstance(self.local_filepath, str):
local_filepath = [self.local_filepath]
else:
diff --git a/providers/sftp/src/airflow/providers/sftp/triggers/sftp.py
b/providers/sftp/src/airflow/providers/sftp/triggers/sftp.py
index a46f29d5a4a..16c894dd3b1 100644
--- a/providers/sftp/src/airflow/providers/sftp/triggers/sftp.py
+++ b/providers/sftp/src/airflow/providers/sftp/triggers/sftp.py
@@ -20,16 +20,29 @@ from __future__ import annotations
import asyncio
from collections.abc import AsyncIterator
from datetime import datetime
+from functools import cached_property
from typing import Any
from dateutil.parser import parse as parse_date
from airflow.providers.common.compat.sdk import AirflowException, timezone
-from airflow.providers.sftp.hooks.sftp import SFTPHookAsync
+from airflow.providers.sftp.hooks.sftp import SFTPHookAsync, SFTPOperation
from airflow.triggers.base import BaseTrigger, TriggerEvent
-class SFTPTrigger(BaseTrigger):
+class BaseSFTPTrigger(BaseTrigger):
+ """Base class for SFTP triggers, providing shared async hook
construction."""
+
+ def __init__(self, sftp_conn_id: str = "sftp_default", remote_host: str |
None = None) -> None:
+ super().__init__()
+ self.sftp_conn_id = sftp_conn_id
+ self.remote_host = remote_host
+
+ def _get_async_hook(self) -> SFTPHookAsync:
+ return SFTPHookAsync(sftp_conn_id=self.sftp_conn_id,
host=self.remote_host)
+
+
+class SFTPTrigger(BaseSFTPTrigger):
"""
SFTPTrigger that fires in below listed scenarios.
@@ -53,10 +66,9 @@ class SFTPTrigger(BaseTrigger):
newer_than: datetime | str | None = None,
poke_interval: float = 5,
) -> None:
- super().__init__()
+ super().__init__(sftp_conn_id=sftp_conn_id)
self.path = path
self.file_pattern = file_pattern
- self.sftp_conn_id = sftp_conn_id
self.newer_than = newer_than
self.poke_interval = poke_interval
@@ -84,30 +96,12 @@ class SFTPTrigger(BaseTrigger):
"""
hook = self._get_async_hook()
- if isinstance(self.newer_than, str):
- self.newer_than = parse_date(self.newer_than)
- _newer_than = timezone.convert_to_utc(self.newer_than) if
self.newer_than else None
while True:
try:
if self.file_pattern:
- files_returned_by_hook = await
hook.get_files_and_attrs_by_pattern(
- path=self.path, fnmatch_pattern=self.file_pattern
+ files_sensed = await hook.sense_files_by_pattern(
+ path=self.path, fnmatch_pattern=self.file_pattern,
newer_than=self.newer_than_utc
)
- files_sensed = []
- for file in files_returned_by_hook:
- if _newer_than:
- if file.attrs.mtime is None:
- continue
- mod_time =
datetime.fromtimestamp(float(file.attrs.mtime)).strftime(
- "%Y%m%d%H%M%S"
- )
- mod_time_utc = timezone.convert_to_utc(
- datetime.strptime(mod_time, "%Y%m%d%H%M%S")
- )
- if _newer_than <= mod_time_utc:
- files_sensed.append(file.filename)
- else:
- files_sensed.append(file.filename)
if files_sensed:
yield TriggerEvent(
{
@@ -116,16 +110,9 @@ class SFTPTrigger(BaseTrigger):
}
)
return
- else:
- mod_time = await hook.get_mod_time(self.path)
- if _newer_than:
- mod_time_utc =
timezone.convert_to_utc(datetime.strptime(mod_time, "%Y%m%d%H%M%S"))
- if _newer_than <= mod_time_utc:
- yield TriggerEvent({"status": "success",
"message": f"Sensed file: {self.path}"})
- return
- else:
- yield TriggerEvent({"status": "success", "message":
f"Sensed file: {self.path}"})
- return
+ elif await hook.sense_path(path=self.path,
newer_than=self.newer_than_utc):
+ yield TriggerEvent({"status": "success", "message":
f"Sensed file: {self.path}"})
+ return
await asyncio.sleep(self.poke_interval)
except AirflowException:
await asyncio.sleep(self.poke_interval)
@@ -138,5 +125,89 @@ class SFTPTrigger(BaseTrigger):
yield TriggerEvent({"status": "error", "message": str(exc)})
- def _get_async_hook(self) -> SFTPHookAsync:
- return SFTPHookAsync(sftp_conn_id=self.sftp_conn_id)
+ @cached_property
+ def newer_than_utc(self) -> datetime | None:
+ """Parse and convert ``newer_than`` to a UTC datetime once, without
mutating the original value."""
+ if not self.newer_than:
+ return None
+ newer_than = parse_date(self.newer_than) if
isinstance(self.newer_than, str) else self.newer_than
+ return timezone.convert_to_utc(newer_than)
+
+
+class SFTPTransferTrigger(BaseSFTPTrigger):
+ """
+ Trigger for SFTPOperator deferrable mode.
+
+ Fires when a file transfer (PUT, GET, or DELETE) completes
+ on the SFTP server, freeing the worker slot during the transfer.
+
+ :param sftp_conn_id: The SFTP connection ID to use.
+ :param local_filepath: Local file path(s) to transfer.
+ :param remote_filepath: Remote file path(s) on the SFTP server.
+ :param operation: The SFTP operation - put, get, or delete.
+ :param confirm: Whether to confirm the file transfer.
+ :param create_intermediate_dirs: Whether to create intermediate dirs.
+ :param remote_host: Remote host to connect to (overrides connection).
+ :param concurrency: Number of threads for directory transfers.
+ :param prefetch: Whether to prefetch during file retrieval.
+ """
+
+ def __init__(
+ self,
+ sftp_conn_id: str = "sftp_default",
+ local_filepath: str | list[str] | None = None,
+ remote_filepath: str | list[str] = "",
+ operation: str = SFTPOperation.PUT,
+ confirm: bool = True,
+ create_intermediate_dirs: bool = False,
+ remote_host: str | None = None,
+ concurrency: int = 1,
+ prefetch: bool = True,
+ ) -> None:
+ super().__init__(sftp_conn_id=sftp_conn_id, remote_host=remote_host)
+ self.local_filepath = local_filepath
+ self.remote_filepath = remote_filepath
+ self.operation = operation
+ self.confirm = confirm
+ self.create_intermediate_dirs = create_intermediate_dirs
+ self.concurrency = concurrency
+ self.prefetch = prefetch
+
+ def serialize(self) -> tuple[str, dict[str, Any]]:
+ """Serialize the trigger for storage in the database."""
+ return (
+ f"{self.__class__.__module__}.{self.__class__.__name__}",
+ {
+ "sftp_conn_id": self.sftp_conn_id,
+ "local_filepath": self.local_filepath,
+ "remote_filepath": self.remote_filepath,
+ "operation": self.operation,
+ "confirm": self.confirm,
+ "create_intermediate_dirs": self.create_intermediate_dirs,
+ "remote_host": self.remote_host,
+ "concurrency": self.concurrency,
+ "prefetch": self.prefetch,
+ },
+ )
+
+ async def run(self) -> AsyncIterator[TriggerEvent]:
+ """Run the file transfer asynchronously and yield a TriggerEvent when
done."""
+ try:
+ hook = self._get_async_hook()
+ await hook.transfer(
+ operation=self.operation,
+ local_filepath=self.local_filepath,
+ remote_filepath=self.remote_filepath,
+ confirm=self.confirm,
+ create_intermediate_dirs=self.create_intermediate_dirs,
+ concurrency=self.concurrency,
+ prefetch=self.prefetch,
+ )
+ yield TriggerEvent(
+ {
+ "status": "success",
+ "local_filepath": self.local_filepath,
+ }
+ )
+ except Exception as e:
+ yield TriggerEvent({"status": "error", "message": str(e)})
diff --git a/providers/sftp/tests/unit/sftp/hooks/test_sftp.py
b/providers/sftp/tests/unit/sftp/hooks/test_sftp.py
index db6d8292299..7b51940172a 100644
--- a/providers/sftp/tests/unit/sftp/hooks/test_sftp.py
+++ b/providers/sftp/tests/unit/sftp/hooks/test_sftp.py
@@ -23,7 +23,8 @@ import os
import shutil
import stat
from io import BytesIO, StringIO
-from unittest.mock import AsyncMock, MagicMock, Mock, PropertyMock, patch
+from types import SimpleNamespace
+from unittest.mock import AsyncMock, MagicMock, Mock, PropertyMock, call, patch
import paramiko
import pytest
@@ -34,7 +35,7 @@ from paramiko.sftp_client import SFTPClient
from airflow.models import Connection
from airflow.providers.common.compat.sdk import AirflowException
-from airflow.providers.sftp.hooks.sftp import SFTPHook, SFTPHookAsync
+from airflow.providers.sftp.hooks.sftp import CHUNK_SIZE, SFTPHook,
SFTPHookAsync, SFTPOperation
def generate_host_key(pkey: paramiko.PKey):
@@ -1131,29 +1132,32 @@ class TestSFTPHookAsync:
sftp_client_mock.__aexit__.assert_awaited()
@pytest.mark.asyncio
- @patch("aiofiles.open")
- async def test_retrieve_file_to_path(self, mock_aiofiles_open,
sftp_hook_mocked):
+ async def test_retrieve_file_to_path(self, sftp_hook_mocked):
"""
- Assert that retrieve_file writes to a local file using aiofiles
+ Assert that retrieve_file downloads to a local path using sftp.get
with concurrent read-ahead requests.
"""
hook, sftp_client_mock = sftp_hook_mocked
sftp_client = sftp_client_mock.__aenter__.return_value
- mock_remote_file = AsyncMock()
- mock_remote_file.read = AsyncMock(side_effect=[b"abc", b"",
StopAsyncIteration])
- sftp_client.open.return_value.__aenter__.return_value =
mock_remote_file
-
- mock_file = AsyncMock()
- aiofiles_cm = AsyncMock()
- aiofiles_cm.__aenter__.return_value = mock_file
- aiofiles_cm.__aexit__.return_value = None
- mock_aiofiles_open.return_value = aiofiles_cm
+ sftp_client.get = AsyncMock()
await hook.retrieve_file("/remote/file", "/local/file")
- sftp_client.open.assert_called_once_with("/remote/file", "rb")
- mock_file.write.assert_awaited()
+ sftp_client.get.assert_awaited_once_with("/remote/file",
"/local/file", block_size=CHUNK_SIZE)
sftp_client_mock.__aexit__.assert_awaited()
+ @pytest.mark.asyncio
+ async def test_retrieve_file_to_path_disables_prefetch(self,
sftp_hook_mocked):
+ """Assert that prefetch=False limits read-ahead to one request in
flight by capping max_requests to 1."""
+ hook, sftp_client_mock = sftp_hook_mocked
+
+ sftp_client = sftp_client_mock.__aenter__.return_value
+ sftp_client.get = AsyncMock()
+
+ await hook.retrieve_file("/remote/file", "/local/file", prefetch=False)
+ sftp_client.get.assert_awaited_once_with(
+ "/remote/file", "/local/file", block_size=CHUNK_SIZE,
max_requests=1
+ )
+
@pytest.mark.asyncio
async def test_retrieve_file_to_bytesio(self, sftp_hook_mocked):
"""
@@ -1179,22 +1183,62 @@ class TestSFTPHookAsync:
sftp_client = sftp_client_mock.__aenter__.return_value
sftp_client.makedirs = AsyncMock()
- await hook.store_file("/remote/new/dir/file.txt", BytesIO(b"abc"))
+ await hook.store_file("/remote/new/dir/file.txt", BytesIO(b"abc"),
confirm=False)
sftp_client.makedirs.assert_awaited_once_with("/remote/new/dir")
sftp_client.open.assert_called_once_with("/remote/new/dir/file.txt",
"wb")
@pytest.mark.asyncio
- async def test_store_file_path_creates_parent_directories(self,
sftp_hook_mocked):
+ async def test_store_file_bytesio_confirms_uploaded_size(self,
sftp_hook_mocked):
+ """Assert that confirm=True checks the remote file size against the
uploaded byte count."""
+ hook, sftp_client_mock = sftp_hook_mocked
+ sftp_client = sftp_client_mock.__aenter__.return_value
+ sftp_client.makedirs = AsyncMock()
+ sftp_client.stat = AsyncMock(return_value=Mock(spec=SFTPAttrs, size=3))
+
+ await hook.store_file("/remote/new/dir/file.txt", BytesIO(b"abc"),
confirm=True)
+
+ sftp_client.stat.assert_awaited_once_with("/remote/new/dir/file.txt")
+
+ @pytest.mark.asyncio
+ async def test_store_file_bytesio_confirm_mismatch_raises(self,
sftp_hook_mocked):
+ """Assert that a size mismatch on confirm raises OSError, mirroring
the sync hook."""
+ hook, sftp_client_mock = sftp_hook_mocked
+ sftp_client = sftp_client_mock.__aenter__.return_value
+ sftp_client.makedirs = AsyncMock()
+ sftp_client.stat = AsyncMock(return_value=Mock(spec=SFTPAttrs, size=1))
+
+ with pytest.raises(OSError, match="size mismatch in put"):
+ await hook.store_file("/remote/new/dir/file.txt", BytesIO(b"abc"),
confirm=True)
+
+ @pytest.mark.asyncio
+ async def test_store_file_path_creates_parent_directories(self,
sftp_hook_mocked, tmp_path):
"""Assert that local-path uploads create parent dirs before put()."""
hook, sftp_client_mock = sftp_hook_mocked
sftp_client = sftp_client_mock.__aenter__.return_value
sftp_client.makedirs = AsyncMock()
- await hook.store_file("/remote/new/dir/file.txt", "/local/file.txt")
+ local_file = tmp_path / "file.txt"
+ local_file.write_bytes(b"abc")
+
+ await hook.store_file("/remote/new/dir/file.txt", str(local_file),
confirm=False)
sftp_client.makedirs.assert_awaited_once_with("/remote/new/dir")
- sftp_client.put.assert_awaited_once_with("/local/file.txt",
"/remote/new/dir/file.txt")
+ sftp_client.put.assert_awaited_once_with(str(local_file),
"/remote/new/dir/file.txt")
+
+ @pytest.mark.asyncio
+ async def test_store_file_path_confirm_mismatch_raises(self,
sftp_hook_mocked, tmp_path):
+ """Assert that confirm=True (the default) verifies size for path-based
uploads too."""
+ hook, sftp_client_mock = sftp_hook_mocked
+ sftp_client = sftp_client_mock.__aenter__.return_value
+ sftp_client.makedirs = AsyncMock()
+ sftp_client.stat = AsyncMock(return_value=Mock(spec=SFTPAttrs, size=1))
+
+ local_file = tmp_path / "file.txt"
+ local_file.write_bytes(b"abc")
+
+ with pytest.raises(OSError, match="size mismatch in put"):
+ await hook.store_file("/remote/new/dir/file.txt", str(local_file))
@pytest.mark.asyncio
async def test_store_file_rejects_raw_bytes(self, sftp_hook_mocked):
@@ -1403,3 +1447,385 @@ class TestSFTPHookAsync:
assert files is not None
assert sorted(files) == sorted(["file1", "subdir/file2"])
sftp_client_mock.__aexit__.assert_awaited()
+
+ @pytest.mark.asyncio
+ async def test_isdir_true_for_directory(self, sftp_hook_mocked):
+ hook, sftp_client_mock = sftp_hook_mocked
+ sftp_client = sftp_client_mock.__aenter__.return_value
+ sftp_client.stat = AsyncMock(return_value=Mock(spec=SFTPAttrs,
permissions=stat.S_IFDIR))
+
+ assert await hook.isdir("/remote/dir") is True
+
+ @pytest.mark.asyncio
+ async def test_isdir_false_for_file(self, sftp_hook_mocked):
+ hook, sftp_client_mock = sftp_hook_mocked
+ sftp_client = sftp_client_mock.__aenter__.return_value
+ sftp_client.stat = AsyncMock(return_value=Mock(spec=SFTPAttrs,
permissions=stat.S_IFREG))
+
+ assert await hook.isdir("/remote/file") is False
+
+ @pytest.mark.asyncio
+ async def test_isdir_false_when_missing(self, sftp_hook_mocked):
+ hook, sftp_client_mock = sftp_hook_mocked
+ sftp_client = sftp_client_mock.__aenter__.return_value
+ sftp_client.stat = AsyncMock(side_effect=SFTPNoSuchFile("no such
file"))
+
+ assert await hook.isdir("/remote/missing") is False
+
+ @pytest.mark.asyncio
+ async def test_path_exists_true(self, sftp_hook_mocked):
+ hook, sftp_client_mock = sftp_hook_mocked
+ sftp_client = sftp_client_mock.__aenter__.return_value
+ sftp_client.stat = AsyncMock(return_value=Mock(spec=SFTPAttrs))
+
+ assert await hook.path_exists("/remote/file") is True
+
+ @pytest.mark.asyncio
+ async def test_path_exists_false_when_missing(self, sftp_hook_mocked):
+ hook, sftp_client_mock = sftp_hook_mocked
+ sftp_client = sftp_client_mock.__aenter__.return_value
+ sftp_client.stat = AsyncMock(side_effect=SFTPNoSuchFile("no such
file"))
+
+ assert await hook.path_exists("/remote/missing") is False
+
+ @pytest.mark.asyncio
+ async def test_create_directory_is_idempotent(self, sftp_hook_mocked):
+ """Assert that create_directory uses exist_ok=True, mirroring
SFTPHook.create_directory."""
+ hook, sftp_client_mock = sftp_hook_mocked
+ sftp_client = sftp_client_mock.__aenter__.return_value
+ sftp_client.makedirs = AsyncMock()
+
+ await hook.create_directory("/remote/new/dir")
+ sftp_client.makedirs.assert_awaited_once_with("/remote/new/dir",
exist_ok=True)
+
+ @pytest.mark.asyncio
+ async def test_delete_file_calls_unlink(self, sftp_hook_mocked):
+ hook, sftp_client_mock = sftp_hook_mocked
+ sftp_client = sftp_client_mock.__aenter__.return_value
+ sftp_client.unlink = AsyncMock()
+
+ await hook.delete_file("/remote/file")
+ sftp_client.unlink.assert_awaited_once_with("/remote/file")
+
+ @pytest.mark.asyncio
+ async def test_delete_directory_without_files_only_removes_directory(self,
sftp_hook_mocked):
+ hook, sftp_client_mock = sftp_hook_mocked
+ sftp_client = sftp_client_mock.__aenter__.return_value
+ sftp_client.remove = AsyncMock()
+ sftp_client.rmdir = AsyncMock()
+
+ await hook.delete_directory("/remote/dir")
+
+ sftp_client.remove.assert_not_awaited()
+ sftp_client.rmdir.assert_awaited_once_with("/remote/dir")
+
+ @pytest.mark.asyncio
+ async def test_delete_directory_include_files_removes_deepest_first(self,
sftp_hook_mocked):
+ hook, sftp_client_mock = sftp_hook_mocked
+ sftp_client = sftp_client_mock.__aenter__.return_value
+ sftp_client.remove = AsyncMock()
+ sftp_client.rmdir = AsyncMock()
+
+ with patch.object(
+ hook,
+ "get_tree_map",
+ AsyncMock(return_value=(["/remote/dir/file1",
"/remote/dir/sub/file2"], ["/remote/dir/sub"], [])),
+ ):
+ await hook.delete_directory("/remote/dir", include_files=True)
+
+ sftp_client.remove.assert_any_await("/remote/dir/file1")
+ sftp_client.remove.assert_any_await("/remote/dir/sub/file2")
+ assert sftp_client.rmdir.await_args_list == [
+ call("/remote/dir/sub"),
+ call("/remote/dir"),
+ ]
+
+ @pytest.mark.asyncio
+ async def test_get_tree_map_splits_files_dirs_and_unknowns(self,
sftp_hook_mocked):
+ hook, sftp_client_mock = sftp_hook_mocked
+ sftp_client = sftp_client_mock.__aenter__.return_value
+
+ async def readdir_side_effect(path):
+ if path == "/dir":
+
+ class File:
+ filename = "file1"
+ attrs = type("attrs", (), {"permissions": stat.S_IFREG})
+
+ class Subdir:
+ filename = "subdir"
+ attrs = type("attrs", (), {"permissions": stat.S_IFDIR})
+
+ class Fifo:
+ filename = "fifo1"
+ attrs = type("attrs", (), {"permissions": stat.S_IFIFO})
+
+ return [File(), Subdir(), Fifo()]
+ return []
+
+ sftp_client.readdir.side_effect = readdir_side_effect
+ sftp_client.realpath.side_effect = lambda path: path
+
+ files, dirs, unknowns = await hook.get_tree_map("/dir")
+ assert files == ["/dir/file1"]
+ assert dirs == ["/dir/subdir"]
+ assert unknowns == ["/dir/fifo1"]
+
+ @pytest.mark.asyncio
+ async def test_retrieve_directory_rejects_existing_local_path(self,
sftp_hook_mocked, tmp_path):
+ hook, _ = sftp_hook_mocked
+
+ with pytest.raises(FileExistsError, match="already exists"):
+ await hook.retrieve_directory("/remote/dir", str(tmp_path))
+
+ @pytest.mark.asyncio
+ async def
test_retrieve_directory_downloads_files_and_creates_subdirs(self,
sftp_hook_mocked, tmp_path):
+ hook, _ = sftp_hook_mocked
+ local_dir = tmp_path / "download"
+
+ with (
+ patch.object(
+ hook,
+ "get_tree_map",
+ AsyncMock(
+ return_value=(
+ ["/remote/dir/file1", "/remote/dir/sub/file2"],
+ ["/remote/dir/sub"],
+ [],
+ )
+ ),
+ ),
+ patch.object(hook, "retrieve_file", AsyncMock()) as
mock_retrieve_file,
+ ):
+ await hook.retrieve_directory("/remote/dir", str(local_dir))
+
+ assert local_dir.is_dir()
+ assert (local_dir / "sub").is_dir()
+ mock_retrieve_file.assert_any_await("/remote/dir/file1", str(local_dir
/ "file1"), prefetch=True)
+ mock_retrieve_file.assert_any_await(
+ "/remote/dir/sub/file2", str(local_dir / "sub" / "file2"),
prefetch=True
+ )
+
+ @pytest.mark.asyncio
+ async def test_retrieve_directory_rejects_path_traversal(self,
sftp_hook_mocked, tmp_path):
+ """Assert that a malicious remote path escaping the local destination
is rejected."""
+ hook, _ = sftp_hook_mocked
+ local_dir = tmp_path / "download"
+
+ with (
+ patch.object(
+ hook,
+ "get_tree_map",
+ AsyncMock(return_value=(["/remote/dir/../../evil"], [], [])),
+ ),
+ patch.object(hook, "retrieve_file", AsyncMock()),
+ ):
+ with pytest.raises(ValueError, match="outside the destination
directory"):
+ await hook.retrieve_directory("/remote/dir", str(local_dir))
+
+ @pytest.mark.asyncio
+ async def test_store_directory_rejects_existing_remote_path(self,
sftp_hook_mocked, tmp_path):
+ hook, _ = sftp_hook_mocked
+
+ with patch.object(hook, "path_exists", AsyncMock(return_value=True)):
+ with pytest.raises(FileExistsError, match="already exists"):
+ await hook.store_directory("/remote/dir", str(tmp_path))
+
+ @pytest.mark.asyncio
+ async def test_store_directory_uploads_tree(self, sftp_hook_mocked,
tmp_path):
+ hook, _ = sftp_hook_mocked
+ local_dir = tmp_path / "upload"
+ (local_dir / "sub").mkdir(parents=True)
+ (local_dir / "file1").write_bytes(b"abc")
+ (local_dir / "sub" / "file2").write_bytes(b"def")
+
+ with (
+ patch.object(hook, "path_exists", AsyncMock(return_value=False)),
+ patch.object(hook, "create_directory", AsyncMock()) as
mock_create_directory,
+ patch.object(hook, "store_file", AsyncMock()) as mock_store_file,
+ ):
+ await hook.store_directory("/remote/dir", str(local_dir))
+
+ mock_create_directory.assert_any_await("/remote/dir")
+ mock_create_directory.assert_any_await(os.path.join("/remote/dir",
"sub"))
+ mock_store_file.assert_any_await(
+ os.path.join("/remote/dir", "file1"), str(local_dir / "file1"),
confirm=True
+ )
+ mock_store_file.assert_any_await(
+ os.path.join("/remote/dir", "sub", "file2"), str(local_dir / "sub"
/ "file2"), confirm=True
+ )
+
+ @pytest.mark.asyncio
+ async def test_transfer_get_dispatches_to_retrieve_directory(self,
sftp_hook_mocked):
+ hook, _ = sftp_hook_mocked
+
+ with (
+ patch.object(hook, "isdir", AsyncMock(return_value=True)),
+ patch.object(hook, "retrieve_directory", AsyncMock()) as
mock_retrieve_directory,
+ patch.object(hook, "retrieve_file", AsyncMock()) as
mock_retrieve_file,
+ ):
+ await hook.transfer(SFTPOperation.GET,
local_filepath="/local/dir", remote_filepath="/remote/dir")
+
+ mock_retrieve_directory.assert_awaited_once_with("/remote/dir",
"/local/dir", prefetch=True)
+ mock_retrieve_file.assert_not_awaited()
+
+ @pytest.mark.asyncio
+ async def test_transfer_get_dispatches_to_retrieve_file(self,
sftp_hook_mocked):
+ hook, _ = sftp_hook_mocked
+
+ with (
+ patch.object(hook, "isdir", AsyncMock(return_value=False)),
+ patch.object(hook, "retrieve_directory", AsyncMock()) as
mock_retrieve_directory,
+ patch.object(hook, "retrieve_file", AsyncMock()) as
mock_retrieve_file,
+ ):
+ await hook.transfer(
+ SFTPOperation.GET, local_filepath="/local/file",
remote_filepath="/remote/file"
+ )
+
+ mock_retrieve_file.assert_awaited_once_with("/remote/file",
"/local/file", prefetch=True)
+ mock_retrieve_directory.assert_not_awaited()
+
+ @pytest.mark.asyncio
+ async def test_transfer_get_creates_intermediate_dirs(self,
sftp_hook_mocked, tmp_path):
+ hook, _ = sftp_hook_mocked
+ local_file = tmp_path / "new" / "dir" / "file"
+
+ with (
+ patch.object(hook, "isdir", AsyncMock(return_value=False)),
+ patch.object(hook, "retrieve_file", AsyncMock()),
+ ):
+ await hook.transfer(
+ SFTPOperation.GET,
+ local_filepath=str(local_file),
+ remote_filepath="/remote/file",
+ create_intermediate_dirs=True,
+ )
+
+ assert local_file.parent.is_dir()
+
+ @pytest.mark.asyncio
+ async def test_transfer_put_dispatches_to_store_directory(self,
sftp_hook_mocked, tmp_path):
+ hook, _ = sftp_hook_mocked
+ local_dir = tmp_path / "dir"
+ local_dir.mkdir()
+
+ with (
+ patch.object(hook, "store_directory", AsyncMock()) as
mock_store_directory,
+ patch.object(hook, "store_file", AsyncMock()) as mock_store_file,
+ ):
+ await hook.transfer(
+ SFTPOperation.PUT, local_filepath=str(local_dir),
remote_filepath="/remote/dir"
+ )
+
+ mock_store_directory.assert_awaited_once_with("/remote/dir",
str(local_dir), confirm=True)
+ mock_store_file.assert_not_awaited()
+
+ @pytest.mark.asyncio
+ async def test_transfer_put_dispatches_to_store_file(self,
sftp_hook_mocked, tmp_path):
+ hook, _ = sftp_hook_mocked
+ local_file = tmp_path / "file"
+ local_file.write_bytes(b"abc")
+
+ with (
+ patch.object(hook, "store_directory", AsyncMock()) as
mock_store_directory,
+ patch.object(hook, "store_file", AsyncMock()) as mock_store_file,
+ ):
+ await hook.transfer(
+ SFTPOperation.PUT, local_filepath=str(local_file),
remote_filepath="/remote/file"
+ )
+
+ mock_store_file.assert_awaited_once_with("/remote/file",
str(local_file), confirm=True)
+ mock_store_directory.assert_not_awaited()
+
+ @pytest.mark.asyncio
+ async def test_transfer_delete_dispatches_to_directory_delete(self,
sftp_hook_mocked):
+ hook, _ = sftp_hook_mocked
+
+ with (
+ patch.object(hook, "isdir", AsyncMock(return_value=True)),
+ patch.object(hook, "delete_directory", AsyncMock()) as
mock_delete_directory,
+ patch.object(hook, "delete_file", AsyncMock()) as mock_delete_file,
+ ):
+ await hook.transfer(SFTPOperation.DELETE, local_filepath=None,
remote_filepath="/remote/dir")
+
+ mock_delete_directory.assert_awaited_once_with("/remote/dir",
include_files=True)
+ mock_delete_file.assert_not_awaited()
+
+ @pytest.mark.asyncio
+ async def test_transfer_delete_missing_file_warns_instead_of_raising(self,
sftp_hook_mocked, caplog):
+ """Assert that deleting a missing file warns and continues, mirroring
SFTPHook.transfer."""
+ hook, _ = sftp_hook_mocked
+
+ with (
+ patch.object(hook, "isdir", AsyncMock(return_value=False)),
+ patch.object(hook, "delete_file",
AsyncMock(side_effect=SFTPNoSuchFile("missing"))),
+ ):
+ await hook.transfer(SFTPOperation.DELETE, local_filepath=None,
remote_filepath="/remote/missing")
+
+ assert "does not exist" in caplog.text
+
+ @pytest.mark.asyncio
+ async def test_transfer_get_directory_uses_single_connection(self,
sftp_hook_mocked, tmp_path):
+ hook, sftp_cm_mock = sftp_hook_mocked
+ sftp_client = sftp_cm_mock.__aenter__.return_value
+ sftp_client.stat =
AsyncMock(return_value=SimpleNamespace(permissions=stat.S_IFDIR))
+ local_dir = tmp_path / "download"
+
+ with (
+ patch.object(
+ hook,
+ "get_tree_map",
+ AsyncMock(return_value=(["/remote/dir/a", "/remote/dir/b",
"/remote/dir/c"], [], [])),
+ ),
+ patch.object(hook, "retrieve_file", AsyncMock()) as
mock_retrieve_file,
+ ):
+ await hook.transfer(
+ SFTPOperation.GET, local_filepath=str(local_dir),
remote_filepath="/remote/dir"
+ )
+
+ assert hook._get_conn.await_count == 1
+ assert sftp_cm_mock.__aexit__.await_count == 1
+ assert mock_retrieve_file.await_count == 3
+
+ @pytest.mark.asyncio
+ async def test_transfer_put_directory_uses_single_connection(self,
sftp_hook_mocked, tmp_path):
+ hook, sftp_cm_mock = sftp_hook_mocked
+ sftp_client = sftp_cm_mock.__aenter__.return_value
+ sftp_client.stat = AsyncMock(side_effect=SFTPNoSuchFile("missing"))
+ local_dir = tmp_path / "upload"
+ (local_dir / "sub").mkdir(parents=True)
+ for name in ("a", "b", "sub/c"):
+ (local_dir / name).write_bytes(b"x")
+
+ with patch.object(hook, "store_file", AsyncMock()) as mock_store_file:
+ await hook.transfer(
+ SFTPOperation.PUT, local_filepath=str(local_dir),
remote_filepath="/remote/dir"
+ )
+
+ assert hook._get_conn.await_count == 1
+ assert sftp_cm_mock.__aexit__.await_count == 1
+ assert mock_store_file.await_count == 3
+ assert sftp_client.makedirs.await_count == 2
+
+ @pytest.mark.asyncio
+ async def
test_get_managed_conn_is_shared_by_nested_calls_and_closed_once(self,
sftp_hook_mocked):
+ hook, sftp_cm_mock = sftp_hook_mocked
+ sftp_client = sftp_cm_mock.__aenter__.return_value
+ sftp_client.stat =
AsyncMock(return_value=SimpleNamespace(permissions=stat.S_IFDIR))
+
+ async with hook.get_managed_conn():
+ assert await hook.isdir("/remote/dir")
+ assert await hook.path_exists("/remote/dir")
+ assert hook.get_conn_count() == 1
+ assert sftp_cm_mock.__aexit__.await_count == 0
+
+ assert hook._get_conn.await_count == 1
+ assert sftp_cm_mock.__aexit__.await_count == 1
+ assert hook.get_conn_count() == 0
+ assert hook.conn is None
+
+
+def test_sftp_operation_values():
+ assert SFTPOperation.GET == "get"
+ assert SFTPOperation.PUT == "put"
+ assert SFTPOperation.DELETE == "delete"
diff --git a/providers/sftp/tests/unit/sftp/operators/test_sftp.py
b/providers/sftp/tests/unit/sftp/operators/test_sftp.py
index e862f8be538..60bec348499 100644
--- a/providers/sftp/tests/unit/sftp/operators/test_sftp.py
+++ b/providers/sftp/tests/unit/sftp/operators/test_sftp.py
@@ -1,4 +1,3 @@
-#
# Licensed to the Apache Software Foundation (ASF) under one
# or more contributor license agreements. See the NOTICE file
# distributed with this work for additional information
@@ -28,11 +27,13 @@ from unittest import mock
import paramiko
import pytest
+from airflow.exceptions import TaskDeferred
from airflow.models import DAG, Connection
from airflow.providers.common.compat.openlineage.facet import Dataset
from airflow.providers.common.compat.sdk import AirflowException, timezone
from airflow.providers.sftp.hooks.sftp import SFTPHook
from airflow.providers.sftp.operators.sftp import SFTPOperation, SFTPOperator
+from airflow.providers.sftp.triggers.sftp import SFTPTransferTrigger
from airflow.providers.ssh.hooks.ssh import SSHHook
from airflow.providers.ssh.operators.ssh import SSHOperator
@@ -383,6 +384,23 @@ class TestSFTPOperator:
remote_filepath=["/tmp/test1", "/tmp/test2"],
).execute(None)
+ @pytest.mark.parametrize(
+ "operation", ["GET", "Get", "get", "PUT", "Put", "put", "DELETE",
"Delete", "delete"]
+ )
+ @mock.patch("airflow.providers.sftp.operators.sftp.SFTPHook.transfer",
autospec=True)
+ def test_operation_is_case_insensitive(self, mock_transfer, operation):
+ """Mixed-case operation values must not raise, matching the
pre-refactor behavior."""
+ local_filepath = None if operation.lower() == SFTPOperation.DELETE
else "/tmp/test"
+ SFTPOperator(
+ task_id="test_sftp_case_insensitive_operation",
+ sftp_hook=self.sftp_hook,
+ local_filepath=local_filepath,
+ remote_filepath="/tmp/remotetest",
+ operation=operation,
+ ).execute(None)
+ assert mock_transfer.call_count == 1
+ assert mock_transfer.call_args.kwargs["operation"] == operation
+
@mock.patch("airflow.providers.sftp.operators.sftp.SFTPHook.retrieve_file")
def test_str_filepaths_get(self, mock_get):
local_filepath = "/tmp/test"
@@ -566,11 +584,18 @@ class TestSFTPOperator:
args, _ = mock_delete.call_args_list[0]
assert args == (remote_filepath,)
+ @pytest.mark.parametrize(
+ "missing_error",
+ [
+ pytest.param(FileNotFoundError("missing"), id="file_not_found"),
+ pytest.param(OSError(errno.ENOENT, "No such file"),
id="paramiko_enoent_ioerror"),
+ ],
+ )
@mock.patch("airflow.providers.sftp.operators.sftp.SFTPHook.delete_file")
@mock.patch("airflow.providers.sftp.operators.sftp.SFTPHook.isdir")
- def test_delete_missing_file_warns(self, mock_isdir, mock_delete, caplog):
+ def test_delete_missing_file_warns(self, mock_isdir, mock_delete,
missing_error, caplog):
mock_isdir.return_value = False
- mock_delete.side_effect = FileNotFoundError("missing")
+ mock_delete.side_effect = missing_error
remote_filepath = "/tmp/missing"
sftp_op = SFTPOperator(
task_id="test_missing_file_delete_warns",
@@ -673,3 +698,183 @@ class TestSFTPOperator:
assert lineage.inputs == expected[0]
assert lineage.outputs == expected[1]
+
+
+class TestSFTPOperatorDeferrable:
+ """Tests for SFTPOperator deferrable mode."""
+
+ def test_sftp_operator_defers_when_deferrable_true(self):
+ """Test that SFTPOperator defers when deferrable=True."""
+ operator = SFTPOperator(
+ task_id="test_sftp_defer",
+ ssh_conn_id="ssh_default",
+ local_filepath="/tmp/test.txt",
+ remote_filepath="/remote/test.txt",
+ operation=SFTPOperation.PUT,
+ deferrable=True,
+ )
+ with pytest.raises(TaskDeferred) as exc:
+ operator.execute(context={})
+ assert isinstance(exc.value.trigger, SFTPTransferTrigger)
+ assert exc.value.method_name == "execute_complete"
+
+ @mock.patch.dict("os.environ", {"AIRFLOW_CONN_MY_PROD_SFTP":
"sftp://[email protected]"})
+ def
test_sftp_operator_defer_uses_sftp_hook_conn_id_when_ssh_conn_id_unset(self):
+ """
+ Assert that deferring honors a supplied sftp_hook's connection id.
+
+ Regression test: previously, when only ``sftp_hook`` (not
``ssh_conn_id``) was
+ provided, the trigger silently fell back to
``SFTPHookAsync.default_conn_name``
+ ("sftp_default") instead of the hook's actual connection, redirecting
the
+ deferred transfer to the wrong server.
+ """
+ operator = SFTPOperator(
+ task_id="test_sftp_defer_hook_conn_id",
+ sftp_hook=SFTPHook(ssh_conn_id="my_prod_sftp"),
+ local_filepath="/tmp/test.txt",
+ remote_filepath="/remote/test.txt",
+ operation=SFTPOperation.PUT,
+ deferrable=True,
+ )
+ with pytest.raises(TaskDeferred) as exc:
+ operator.execute(context={})
+ assert exc.value.trigger.sftp_conn_id == "my_prod_sftp"
+
+ def test_sftp_operator_defer_without_any_conn_id_raises(self):
+ operator = SFTPOperator(
+ task_id="test_sftp_defer_no_conn_id",
+ sftp_hook=SFTPHook(ssh_conn_id=None, remote_host="example.com"),
+ local_filepath="/tmp/test.txt",
+ remote_filepath="/remote/test.txt",
+ operation=SFTPOperation.PUT,
+ deferrable=True,
+ )
+ with pytest.raises(ValueError, match="requires a connection id"):
+ operator.execute(context={})
+
+ def test_sftp_operator_execute_complete_success(self):
+ """Test execute_complete returns local_filepath on success."""
+ operator = SFTPOperator(
+ task_id="test_sftp_complete",
+ ssh_conn_id="ssh_default",
+ local_filepath="/tmp/test.txt",
+ remote_filepath="/remote/test.txt",
+ operation=SFTPOperation.PUT,
+ deferrable=True,
+ )
+ event = {"status": "success", "local_filepath": "/tmp/test.txt"}
+ result = operator.execute_complete(context={}, event=event)
+ assert result == "/tmp/test.txt"
+
+ def test_sftp_operator_execute_complete_raises_on_error(self):
+ """Test execute_complete raises AirflowException on error."""
+ operator = SFTPOperator(
+ task_id="test_sftp_error",
+ ssh_conn_id="ssh_default",
+ local_filepath="/tmp/test.txt",
+ remote_filepath="/remote/test.txt",
+ operation=SFTPOperation.PUT,
+ deferrable=True,
+ )
+ event = {"status": "error", "message": "Connection refused"}
+ with pytest.raises(AirflowException, match="Connection refused"):
+ operator.execute_complete(context={}, event=event)
+
+
+class TestSFTPTransferTrigger:
+ """Tests for SFTPTransferTrigger."""
+
+ def test_serialize_roundtrip(self):
+ """Test that serialize() produces correct output for reconstruction."""
+ trigger = SFTPTransferTrigger(
+ sftp_conn_id="ssh_default",
+ local_filepath="/tmp/test.txt",
+ remote_filepath="/remote/test.txt",
+ operation="put",
+ confirm=True,
+ create_intermediate_dirs=False,
+ remote_host=None,
+ concurrency=1,
+ prefetch=True,
+ )
+ classpath, kwargs = trigger.serialize()
+ assert classpath ==
"airflow.providers.sftp.triggers.sftp.SFTPTransferTrigger"
+ assert kwargs["sftp_conn_id"] == "ssh_default"
+ assert kwargs["local_filepath"] == "/tmp/test.txt"
+ assert kwargs["remote_filepath"] == "/remote/test.txt"
+ assert kwargs["operation"] == "put"
+ assert kwargs["confirm"] is True
+ assert kwargs["remote_host"] is None
+ assert kwargs["concurrency"] == 1
+ assert kwargs["prefetch"] is True
+
+ @mock.patch("airflow.providers.sftp.triggers.sftp.SFTPHookAsync",
autospec=True)
+ def test_get_async_hook_forwards_remote_host(self, mock_hook_async):
+ """Test that an explicit remote_host override reaches SFTPHookAsync,
not just the conn_id."""
+ trigger = SFTPTransferTrigger(
+ sftp_conn_id="ssh_default",
+ local_filepath="/tmp/test.txt",
+ remote_filepath="/remote/test.txt",
+ operation="put",
+ remote_host="explicit-host.example.com",
+ )
+ trigger._get_async_hook()
+ mock_hook_async.assert_called_once_with(sftp_conn_id="ssh_default",
host="explicit-host.example.com")
+
+ @mock.patch("airflow.providers.sftp.triggers.sftp.SFTPHookAsync",
autospec=True)
+ def test_get_async_hook_defaults_remote_host_to_none(self,
mock_hook_async):
+ """Test that omitting remote_host does not force an unexpected host
onto the hook."""
+ trigger = SFTPTransferTrigger(
+ sftp_conn_id="ssh_default",
+ local_filepath="/tmp/test.txt",
+ remote_filepath="/remote/test.txt",
+ operation="put",
+ )
+ trigger._get_async_hook()
+ mock_hook_async.assert_called_once_with(sftp_conn_id="ssh_default",
host=None)
+
+ def test_run_success(self):
+ """Test run() yields TriggerEvent with status success."""
+ import asyncio
+
+ trigger = SFTPTransferTrigger(
+ sftp_conn_id="ssh_default",
+ local_filepath="/tmp/test.txt",
+ remote_filepath="/remote/test.txt",
+ operation="put",
+ )
+ with
mock.patch("airflow.providers.sftp.triggers.sftp.SFTPHookAsync.transfer",
return_value=None):
+ events = []
+
+ async def collect():
+ async for event in trigger.run():
+ events.append(event)
+
+ asyncio.run(collect())
+ assert len(events) == 1
+ assert events[0].payload["status"] == "success"
+
+ def test_run_error(self):
+ """Test run() yields TriggerEvent with status error on exception."""
+ import asyncio
+
+ trigger = SFTPTransferTrigger(
+ sftp_conn_id="ssh_default",
+ local_filepath="/tmp/test.txt",
+ remote_filepath="/remote/test.txt",
+ operation="put",
+ )
+ with mock.patch(
+ "airflow.providers.sftp.triggers.sftp.SFTPHookAsync.transfer",
+ side_effect=Exception("Connection failed"),
+ ):
+ events = []
+
+ async def collect():
+ async for event in trigger.run():
+ events.append(event)
+
+ asyncio.run(collect())
+ assert len(events) == 1
+ assert events[0].payload["status"] == "error"
+ assert "Connection failed" in events[0].payload["message"]
diff --git a/providers/sftp/tests/unit/sftp/triggers/test_sftp.py
b/providers/sftp/tests/unit/sftp/triggers/test_sftp.py
index e6c2502f780..ddefa8468a1 100644
--- a/providers/sftp/tests/unit/sftp/triggers/test_sftp.py
+++ b/providers/sftp/tests/unit/sftp/triggers/test_sftp.py
@@ -26,7 +26,7 @@ from unittest import mock
import pytest
from asyncssh.sftp import SFTPAttrs, SFTPName
-from airflow.providers.common.compat.sdk import AirflowException
+from airflow.providers.common.compat.sdk import AirflowException, timezone
from airflow.providers.sftp.triggers.sftp import SFTPTrigger
from airflow.triggers.base import TriggerEvent
@@ -68,6 +68,30 @@ class TestSFTPTrigger:
"poke_interval": 5.0,
}
+ def test_newer_than_utc_none_when_not_provided(self):
+ trigger = SFTPTrigger(path="test/path/")
+ assert trigger.newer_than_utc is None
+
+ def test_newer_than_utc_converts_datetime(self):
+ naive = datetime.datetime(2023, 5, 1, 12, 0, 0)
+ trigger = SFTPTrigger(path="test/path/", newer_than=naive)
+ assert trigger.newer_than_utc == timezone.convert_to_utc(naive)
+
+ def test_newer_than_utc_parses_string(self):
+ trigger = SFTPTrigger(path="test/path/",
newer_than="2023-05-01T12:00:00")
+ assert trigger.newer_than_utc ==
timezone.convert_to_utc(datetime.datetime(2023, 5, 1, 12, 0, 0))
+
+ def test_newer_than_utc_is_cached_and_does_not_mutate_newer_than(self):
+ trigger = SFTPTrigger(path="test/path/",
newer_than="2023-05-01T12:00:00")
+
+ first_access = trigger.newer_than_utc
+ # newer_than stays the original string; only the cached property is
converted.
+ assert trigger.newer_than == "2023-05-01T12:00:00"
+
+ trigger.newer_than = "2024-01-01T00:00:00"
+ # Property is cached, so mutating newer_than afterwards has no effect
on it.
+ assert trigger.newer_than_utc is first_access
+
@pytest.mark.asyncio
@pytest.mark.parametrize(
"newer_than",
diff --git a/uv.lock b/uv.lock
index 88a620d2ded..f311a4f04c2 100644
--- a/uv.lock
+++ b/uv.lock
@@ -7932,7 +7932,6 @@ name = "apache-airflow-providers-sftp"
version = "6.0.1"
source = { editable = "providers/sftp" }
dependencies = [
- { name = "aiofiles" },
{ name = "apache-airflow" },
{ name = "apache-airflow-providers-common-compat" },
{ name = "apache-airflow-providers-ssh" },
@@ -7963,7 +7962,6 @@ docs = [
[package.metadata]
requires-dist = [
- { name = "aiofiles", specifier = ">=23.2.0" },
{ name = "apache-airflow", editable = "." },
{ name = "apache-airflow-providers-common-compat", editable =
"providers/common/compat" },
{ name = "apache-airflow-providers-openlineage", marker = "extra ==
'openlineage'", editable = "providers/openlineage" },