# Copyright 2022 The PyAthena authors
#
# Licensed under the MIT License.
# See LICENSE or https://opensource.org/licenses/MIT.
#
# SPDX-License-Identifier: MIT
"""Asynchronous fsspec filesystem for Amazon S3 built on ``S3FileSystem``."""
from __future__ import annotations
import asyncio
import logging
import mimetypes
import os
from multiprocessing import cpu_count
from typing import TYPE_CHECKING, Any, cast
from fsspec.asyn import AsyncFileSystem
from fsspec.callbacks import _DEFAULT_CALLBACK
from pyathena.filesystem.s3 import S3File, S3FileSystem
from pyathena.filesystem.s3_executor import S3AioExecutor, S3Executor, S3ThreadPoolExecutor
from pyathena.filesystem.s3_object import (
S3Metadata,
S3MultipartUpload,
S3Object,
S3ObjectVersion,
)
if TYPE_CHECKING:
from datetime import datetime
from pyathena.connection import Connection
_logger = logging.getLogger(__name__)
[docs]
class AioS3FileSystem(AsyncFileSystem):
"""An async filesystem interface for Amazon S3 using fsspec's AsyncFileSystem.
This class wraps ``S3FileSystem`` to provide native asyncio support. Instead of
using ``ThreadPoolExecutor`` for parallel operations, it uses ``asyncio.gather``
with ``asyncio.to_thread`` for natural integration with the asyncio event loop.
The implementation uses composition: an internal ``S3FileSystem`` instance handles
all boto3 calls, while this class delegates to it via ``asyncio.to_thread()``.
This avoids diamond inheritance issues and keeps all boto3 logic in one place.
File handles created by ``_open`` use ``S3AioExecutor`` so that parallel
operations (range reads, multipart uploads) are dispatched through the event
loop with ``asyncio.to_thread`` instead of a ``ThreadPoolExecutor`` per file.
An instance created with ``asynchronous=True`` has no event loop of its own,
so its file handles use a ``ThreadPoolExecutor``.
Attributes:
_sync_fs: The internal synchronous S3FileSystem instance.
Example:
>>> from pyathena.filesystem.s3_async import AioS3FileSystem
>>> fs = AioS3FileSystem(asynchronous=True)
>>>
>>> # Use in async context
>>> files = await fs._ls('s3://my-bucket/data/')
>>>
>>> # Sync wrappers (auto-generated by fsspec) need an instance created
>>> # without asynchronous=True; they block the caller until done
>>> files = AioS3FileSystem().ls('s3://my-bucket/data/')
"""
# https://docs.aws.amazon.com/AmazonS3/latest/API/API_DeleteObjects.html
DELETE_OBJECTS_MAX_KEYS: int = 1000
protocol = ("s3", "s3a")
mirror_sync_methods = True
async_impl = True
_extra_tokenize_attributes = ("default_block_size",)
[docs]
def __init__(
self,
connection: Connection[Any] | None = None,
default_block_size: int | None = None,
default_cache_type: str | None = None,
max_workers: int = (cpu_count() or 1) * 5,
s3_additional_kwargs: dict[str, Any] | None = None,
allow_bucket_creation: bool = False,
allow_bucket_deletion: bool = False,
version_aware: bool = False,
asynchronous: bool = False,
loop: Any | None = None,
batch_size: int | None = None,
**kwargs,
) -> None:
"""Initialize the filesystem and its internal ``S3FileSystem``.
Args:
connection: Passed to the internal ``S3FileSystem``.
default_block_size: Passed to the internal ``S3FileSystem``.
default_cache_type: Passed to the internal ``S3FileSystem``.
max_workers: Passed to the internal ``S3FileSystem``.
s3_additional_kwargs: Passed to the internal ``S3FileSystem``.
allow_bucket_creation: Passed to the internal ``S3FileSystem``.
allow_bucket_deletion: Passed to the internal ``S3FileSystem``.
version_aware: Passed to the internal ``S3FileSystem``.
asynchronous: Passed to ``fsspec.asyn.AsyncFileSystem``.
loop: Passed to ``fsspec.asyn.AsyncFileSystem``.
batch_size: Passed to ``fsspec.asyn.AsyncFileSystem``.
**kwargs: Passed to both ``fsspec.asyn.AsyncFileSystem`` and the
internal ``S3FileSystem``.
"""
super().__init__(
asynchronous=asynchronous,
loop=loop,
batch_size=batch_size,
**kwargs,
)
self._sync_fs = S3FileSystem(
connection=connection,
default_block_size=default_block_size,
default_cache_type=default_cache_type,
max_workers=max_workers,
s3_additional_kwargs=s3_additional_kwargs,
allow_bucket_creation=allow_bucket_creation,
allow_bucket_deletion=allow_bucket_deletion,
version_aware=version_aware,
# fsspec caches the AioS3FileSystem itself when caching is wanted.
skip_instance_cache=True,
**kwargs,
)
# Share dircache for cache coherence between async and sync instances
self.dircache = self._sync_fs.dircache
[docs]
@staticmethod
def parse_path(path: str) -> tuple[str, str | None, str | None]:
"""Parse an S3 path into its bucket, key and version ID.
See :meth:`S3FileSystem.parse_path`.
Args:
path: The S3 path.
Returns:
Tuple of the bucket, the key and the version ID.
Raises:
ValueError: If the path is not a valid S3 path.
"""
return S3FileSystem.parse_path(path)
async def _info(self, path: str, **kwargs) -> S3Object:
return await asyncio.to_thread(self._sync_fs.info, path, **kwargs)
async def _ls(self, path: str, detail: bool = False, **kwargs) -> list[S3Object] | list[str]:
return await asyncio.to_thread(self._sync_fs.ls, path, detail=detail, **kwargs)
async def _cat_file(
self, path: str, start: int | None = None, end: int | None = None, **kwargs
) -> bytes:
return await asyncio.to_thread(self._sync_fs.cat_file, path, start=start, end=end, **kwargs)
async def _exists(self, path: str, **kwargs) -> bool:
return await asyncio.to_thread(self._sync_fs.exists, path, **kwargs)
async def _rm_file(self, path: str, **kwargs) -> None:
await asyncio.to_thread(self._sync_fs.rm_file, path, **kwargs)
async def _pipe_file(
self, path: str, value: bytes | bytearray | memoryview, mode: str = "overwrite", **kwargs
) -> None:
if self._intrans:
# The transaction belongs to this filesystem, not to the internal
# S3FileSystem, so write through open() to defer the commit to it.
await asyncio.to_thread(self._pipe_file_in_transaction, path, value, mode, **kwargs)
return
await asyncio.to_thread(self._sync_fs.pipe_file, path, value, mode=mode, **kwargs)
def _pipe_file_in_transaction(
self, path: str, value: bytes | bytearray | memoryview, mode: str, **kwargs
) -> None:
"""Write bytes into the path as a file of this filesystem's transaction.
Args:
path: S3 path (s3://bucket/key) to write to.
value: The bytes to write.
mode: "overwrite" or "create". With "create", the file is
opened in ``xb`` mode: raise FileExistsError when the object
already exists, including one created before the
transaction is committed, which is not replaced.
**kwargs: Additional parameters passed to ``open()``.
Raises:
FileExistsError: If the mode is "create" and the path already
exists.
ValueError: If the data takes more than
``MULTIPART_UPLOAD_MAX_PARTS`` blocks.
"""
block_size = kwargs.get("block_size") or self._sync_fs.default_block_size
# The size in bytes; the length of a memoryview counts its items.
self._sync_fs._check_multipart_upload_size(path, memoryview(value).nbytes, block_size)
self._sync_fs._write_and_close(
self.open(path, "xb" if mode == "create" else "wb", **kwargs), value
)
async def _put_file(
self,
lpath: str,
rpath: str,
callback=_DEFAULT_CALLBACK,
mode: str = "overwrite",
**kwargs,
) -> None:
if self._intrans:
# See _pipe_file.
await asyncio.to_thread(
self._put_file_in_transaction, lpath, rpath, callback, mode, **kwargs
)
return
await asyncio.to_thread(
self._sync_fs.put_file, lpath, rpath, callback=callback, mode=mode, **kwargs
)
def _put_file_in_transaction(
self, lpath: str, rpath: str, callback, mode: str, **kwargs
) -> None:
"""Upload a local file as a file of this filesystem's transaction.
Mirrors :meth:`S3FileSystem.put_file`, but writes through ``open()``
of this filesystem.
Args:
lpath: Local file path to upload.
rpath: S3 destination path (s3://bucket/key).
callback: Progress callback for tracking upload progress.
mode: "overwrite" or "create". With "create", the file is
opened in ``xb`` mode: raise FileExistsError when the object
already exists, including one created before the
transaction is committed, which is not replaced.
**kwargs: Additional S3 parameters (e.g., ContentType, StorageClass).
The ``block_size``, ``max_workers``, and ``s3_additional_kwargs``
parameters of ``open()`` are also accepted.
Raises:
FileExistsError: If the mode is "create" and the path already
exists.
ValueError: If the file takes more than
``MULTIPART_UPLOAD_MAX_PARTS`` blocks.
"""
if os.path.isdir(lpath):
return
_, key, _ = self.parse_path(rpath)
if not key:
return
size = os.path.getsize(lpath)
block_size = kwargs.pop("block_size", None) or self._sync_fs.default_block_size
max_workers = kwargs.pop("max_workers", self._sync_fs.max_workers)
# The other parameters are S3 request parameters, as in pipe_file().
s3_additional_kwargs = {**kwargs.pop("s3_additional_kwargs", {}), **kwargs}
self._sync_fs._check_multipart_upload_size(rpath, size, block_size)
callback.set_size(size)
if "ContentType" not in {**self._sync_fs.s3_additional_kwargs, **s3_additional_kwargs}:
content_type, _ = mimetypes.guess_type(lpath)
if content_type is not None:
s3_additional_kwargs["ContentType"] = content_type
with (
self.open(
rpath,
"xb" if mode == "create" else "wb",
block_size=block_size,
max_workers=max_workers,
s3_additional_kwargs=s3_additional_kwargs,
) as remote,
open(lpath, "rb") as local,
):
while data := local.read(remote.blocksize):
remote.write(data)
callback.relative_update(len(data))
self.invalidate_cache(rpath)
async def _get_file(self, rpath: str, lpath: str, callback=_DEFAULT_CALLBACK, **kwargs) -> None:
await asyncio.to_thread(self._sync_fs.get_file, rpath, lpath, callback=callback, **kwargs)
async def _mkdir(self, path: str, create_parents: bool = True, **kwargs) -> None:
await asyncio.to_thread(self._sync_fs.mkdir, path, create_parents=create_parents, **kwargs)
async def _makedirs(self, path: str, exist_ok: bool = False) -> None:
await asyncio.to_thread(self._sync_fs.makedirs, path, exist_ok=exist_ok)
async def _rm(
self,
path: str | list[str],
recursive: bool = False,
batch_size: int | None = None,
maxdepth: int | None = None,
**kwargs,
) -> None:
"""Delete objects with DeleteObjects requests.
See :meth:`S3FileSystem.rm`. The requests run in parallel with
``asyncio.gather`` and ``asyncio.to_thread``.
Args:
path: S3 path (s3://bucket/key) or list of paths to delete.
recursive: Whether to delete all objects below the paths.
batch_size: Accepted for fsspec compatibility; not used.
maxdepth: Maximum depth to expand when ``recursive`` is True.
**kwargs: Additional parameters passed to the DeleteObjects API.
``Quiet`` (default True) sets the quiet mode of the requests.
Raises:
ValueError: If a path is a bucket.
OSError: If S3 could not delete some of the objects.
"""
paths = await asyncio.to_thread(
self._sync_fs._expand_delete_paths, path, recursive=recursive, maxdepth=maxdepth
)
requests = self._sync_fs._delete_objects_requests(paths, **kwargs)
results = await asyncio.gather(
*[
asyncio.to_thread(self._sync_fs._delete_objects_request, request)
for request in requests
],
return_exceptions=True,
)
self._sync_fs._raise_delete_objects_errors(requests, results)
async def _cp_file(self, path1: str, path2: str, **kwargs) -> None:
"""Copy an S3 object, using async parallel multipart upload for large files."""
# fsspec < 2026.6.0 leaks the typo'd "onerror" keyword from mv();
# see S3FileSystem.cp_file.
kwargs.pop("onerror", None)
# Parameters of the multipart copy, not of the S3 requests.
block_size = kwargs.pop("block_size", None)
max_workers = kwargs.pop("max_workers", None)
bucket1, key1, version_id1 = self.parse_path(path1)
bucket2, key2, version_id2 = self.parse_path(path2)
if version_id2:
raise ValueError("Cannot copy to a versioned file.")
if not key1 or not key2:
raise ValueError("Cannot copy buckets.")
info1 = await self._info(path1)
size1 = info1.get("size", 0)
if size1 <= S3FileSystem.MULTIPART_UPLOAD_MAX_PART_SIZE:
await asyncio.to_thread(
self._sync_fs._copy_object,
bucket1=bucket1,
key1=key1,
version_id1=version_id1,
bucket2=bucket2,
key2=key2,
**kwargs,
)
else:
await self._copy_object_with_multipart_upload(
bucket1=bucket1,
key1=key1,
version_id1=version_id1,
size1=size1,
bucket2=bucket2,
key2=key2,
max_workers=max_workers,
block_size=block_size,
**kwargs,
)
self._sync_fs.invalidate_cache(path2)
async def _copy_object_with_multipart_upload(
self,
bucket1: str,
key1: str,
size1: int,
bucket2: str,
key2: str,
max_workers: int | None = None,
block_size: int | None = None,
version_id1: str | None = None,
**kwargs,
) -> None:
max_workers = max_workers if max_workers else self._sync_fs.max_workers
block_size = block_size if block_size else S3FileSystem.MULTIPART_UPLOAD_MAX_PART_SIZE
if (
block_size < S3FileSystem.MULTIPART_UPLOAD_MIN_PART_SIZE
or block_size > S3FileSystem.MULTIPART_UPLOAD_MAX_PART_SIZE
):
raise ValueError(
"Block size must be between "
f"5 MiB ({S3FileSystem.MULTIPART_UPLOAD_MIN_PART_SIZE} bytes) and "
f"5 GiB ({S3FileSystem.MULTIPART_UPLOAD_MAX_PART_SIZE} bytes), "
f"inclusive: {block_size}."
)
copy_source: dict[str, Any] = {
"Bucket": bucket1,
"Key": key1,
}
if version_id1:
copy_source["VersionId"] = version_id1
ranges = self._sync_fs._get_copy_ranges(size1, block_size)
multipart_upload = await asyncio.to_thread(
self._sync_fs._create_multipart_upload,
bucket=bucket2,
key=key2,
**kwargs,
)
semaphore = asyncio.Semaphore(max_workers)
part_kwargs = self._sync_fs._get_operation_kwargs("upload_part_copy", kwargs)
async def _upload_part(i: int, range_: tuple[int, int]) -> dict[str, Any]:
async with semaphore:
result = await asyncio.to_thread(
self._sync_fs._upload_part_copy,
bucket=bucket2,
key=key2,
copy_source=copy_source,
upload_id=cast(str, multipart_upload.upload_id),
part_number=i + 1,
copy_source_ranges=range_,
**part_kwargs,
)
return {
"ETag": result.etag,
"PartNumber": result.part_number,
}
parts = await asyncio.gather(*[_upload_part(i, r) for i, r in enumerate(ranges)])
parts_list = sorted(parts, key=lambda x: x["PartNumber"])
await asyncio.to_thread(
self._sync_fs._complete_multipart_upload,
bucket=bucket2,
key=key2,
upload_id=cast(str, multipart_upload.upload_id),
parts=parts_list,
**self._sync_fs._get_operation_kwargs("complete_multipart_upload", kwargs),
)
async def _find(
self,
path: str,
maxdepth: int | None = None,
withdirs: bool = False,
**kwargs,
) -> dict[str, S3Object] | list[str]:
detail = kwargs.pop("detail", False)
files = await asyncio.to_thread(
self._sync_fs._find, path, maxdepth=maxdepth, withdirs=withdirs, **kwargs
)
if detail:
return {f.name: f for f in files}
return [f.name for f in files]
def _create_executor(self, max_workers: int) -> S3Executor:
"""Create the executor for the parallel operations of a file.
An instance created with ``asynchronous=True`` has no event loop of
its own, so its files run the operations in a thread pool.
Args:
max_workers: The maximum number of operations that run at once.
Returns:
An ``S3AioExecutor`` on the event loop of this filesystem, or an
``S3ThreadPoolExecutor`` if it has none.
"""
if self._loop is None:
return S3ThreadPoolExecutor(max_workers=max_workers)
return S3AioExecutor(loop=self._loop, max_workers=max_workers)
def _open(
self,
path: str,
mode: str = "rb",
block_size: int | None = None,
cache_type: str | None = None,
autocommit: bool = True,
cache_options: dict[Any, Any] | None = None,
**kwargs,
) -> AioS3File:
if block_size is None:
block_size = self._sync_fs.default_block_size
if cache_type is None:
cache_type = self._sync_fs.default_cache_type
max_workers = kwargs.pop("max_workers", self._sync_fs.max_workers)
# The parameters of the call take precedence over those of the
# filesystem; the caller's dictionary is not modified.
s3_additional_kwargs = {
**self._sync_fs.s3_additional_kwargs,
**kwargs.pop("s3_additional_kwargs", {}),
}
return AioS3File(
self._sync_fs,
path,
mode,
max_workers=max_workers,
executor=self._create_executor(max_workers=max_workers),
block_size=block_size,
cache_type=cache_type,
autocommit=autocommit,
cache_options=cache_options,
s3_additional_kwargs=s3_additional_kwargs,
**kwargs,
)
async def _rmdir(self, path: str) -> None:
await asyncio.to_thread(self._sync_fs.rmdir, path)
[docs]
def rmdir(self, path: str) -> None:
"""Remove an S3 bucket, which must be empty.
See :meth:`S3FileSystem.rmdir`.
Args:
path: S3 bucket path (e.g., "s3://bucket").
"""
self._sync_fs.rmdir(path)
[docs]
def sign(self, path: str, expiration: int = 3600, **kwargs) -> str:
"""Generate a presigned URL for S3 object access.
See :meth:`S3FileSystem.sign`.
Args:
path: S3 path (s3://bucket/key) to generate the URL for.
expiration: URL expiration time in seconds.
**kwargs: Additional parameters passed to :meth:`S3FileSystem.sign`.
Returns:
The presigned URL.
"""
return cast(str, self._sync_fs.sign(path, expiration=expiration, **kwargs))
[docs]
def getxattr(self, path: str, attr_name: str, **kwargs) -> str | None:
"""Get an attribute from the user-defined metadata of the path.
See :meth:`S3FileSystem.getxattr`.
Args:
path: S3 path (s3://bucket/key) to get the attribute for.
attr_name: The name of the attribute.
**kwargs: Additional parameters passed to the HeadObject API.
Returns:
The value of the attribute, or None if the attribute is not set.
"""
return self._sync_fs.getxattr(path, attr_name, **kwargs)
[docs]
def setxattr(self, path: str, copy_kwargs: dict[str, Any] | None = None, **kwargs) -> None:
"""Set the user-defined metadata of the path.
See :meth:`S3FileSystem.setxattr`.
Args:
path: S3 path (s3://bucket/key) to set metadata for.
copy_kwargs: Additional parameters to use for the underlying
CopyObject API call.
**kwargs: Key-value pairs of metadata to set; a None value deletes the
key.
"""
self._sync_fs.setxattr(path, copy_kwargs=copy_kwargs, **kwargs)
[docs]
def chmod(self, path: str, acl: str, recursive: bool = False, **kwargs) -> None:
"""Set the canned ACL of a bucket or key.
See :meth:`S3FileSystem.chmod`.
Args:
path: S3 path (s3://bucket or s3://bucket/key) to set the ACL on.
acl: The canned ACL to apply.
recursive: Whether to apply the ACL to all keys below the path too.
**kwargs: Additional parameters passed to the PutObjectAcl or
PutBucketAcl API.
"""
self._sync_fs.chmod(path, acl, recursive=recursive, **kwargs)
[docs]
def object_version_info(
self, path: str, delete_markers: bool = False, **kwargs
) -> list[S3ObjectVersion]:
"""List the versions of the object or of the objects under the path.
See :meth:`S3FileSystem.object_version_info`.
Args:
path: S3 path (s3://bucket/key or a key prefix) to list the versions for.
delete_markers: Whether to include delete markers in the result.
**kwargs: Additional parameters passed to the ListObjectVersions API.
Returns:
List of S3ObjectVersion instances describing the versions.
"""
return self._sync_fs.object_version_info(path, delete_markers=delete_markers, **kwargs)
[docs]
def list_multipart_uploads(self, path: str) -> list[S3MultipartUpload]:
"""List in-progress (incomplete) multipart uploads in a bucket.
See :meth:`S3FileSystem.list_multipart_uploads`.
Args:
path: S3 bucket or key path (e.g., "s3://bucket" or "s3://bucket/prefix").
Returns:
List of S3MultipartUpload instances describing the uploads.
"""
return self._sync_fs.list_multipart_uploads(path)
[docs]
def clear_multipart_uploads(self, path: str) -> None:
"""Abort any incomplete multipart uploads in the bucket.
See :meth:`S3FileSystem.clear_multipart_uploads`.
Args:
path: S3 bucket or key path (e.g., "s3://bucket" or "s3://bucket/prefix").
"""
self._sync_fs.clear_multipart_uploads(path)
[docs]
def checksum(self, path: str, **kwargs) -> int:
"""Get the checksum of an S3 object or directory.
See :meth:`S3FileSystem.checksum`.
Args:
path: S3 path (s3://bucket/key) to get the checksum for.
**kwargs: Additional arguments passed to :meth:`S3FileSystem.checksum`.
Returns:
Integer checksum derived from the ETag or the directory token.
"""
return cast(int, self._sync_fs.checksum(path, **kwargs))
[docs]
def created(self, path: str) -> datetime:
"""Return the creation time of the path.
See :meth:`S3FileSystem.created`.
Args:
path: S3 path (s3://bucket/key).
Returns:
The last-modified time of the object.
"""
return self._sync_fs.created(path)
[docs]
def modified(self, path: str) -> datetime:
"""Return the last-modified time of the path.
See :meth:`S3FileSystem.modified`.
Args:
path: S3 path (s3://bucket/key).
Returns:
The last-modified time of the object.
"""
return self._sync_fs.modified(path)
[docs]
def invalidate_cache(self, path: str | None = None) -> None:
"""Remove the cached entries of the path and its parent paths.
See :meth:`S3FileSystem.invalidate_cache`.
Args:
path: The path to invalidate. If None, clear the whole cache.
"""
self._sync_fs.invalidate_cache(path)
async def _touch(self, path: str, truncate: bool = True, **kwargs) -> dict[str, Any]:
return await asyncio.to_thread(self._sync_fs.touch, path, truncate=truncate, **kwargs)
[docs]
def touch(self, path: str, truncate: bool = True, **kwargs) -> dict[str, Any]:
"""Create an empty object with PutObject.
See :meth:`S3FileSystem.touch`.
Args:
path: S3 path (s3://bucket/key) of the object.
truncate: If True, replace an existing object with an empty one;
if False, raise if the object exists.
**kwargs: Additional parameters passed to the PutObject API.
Returns:
The PutObject response as a dictionary.
"""
return self._sync_fs.touch(path, truncate=truncate, **kwargs)
[docs]
class AioS3File(S3File):
"""Async-aware S3 file handle using ``S3AioExecutor``.
Functionally identical to ``S3File``; exists as a distinct type for
``isinstance`` checks and to document the async execution model.
All parallel operations (range reads, multipart uploads) are dispatched
through the ``S3Executor`` interface — the ``S3AioExecutor``
provided by ``AioS3FileSystem`` dispatches them through the event loop with
``asyncio.to_thread`` instead of a ``ThreadPoolExecutor`` per file.
For an ``AioS3FileSystem`` created with ``asynchronous=True``, it is an
``S3ThreadPoolExecutor``.
"""