Source code for multistorageclient.providers.s3

   1# SPDX-FileCopyrightText: Copyright (c) 2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
   2# SPDX-License-Identifier: Apache-2.0
   3#
   4# Licensed under the Apache License, Version 2.0 (the "License");
   5# you may not use this file except in compliance with the License.
   6# You may obtain a copy of the License at
   7#
   8# http://www.apache.org/licenses/LICENSE-2.0
   9#
  10# Unless required by applicable law or agreed to in writing, software
  11# distributed under the License is distributed on an "AS IS" BASIS,
  12# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
  13# See the License for the specific language governing permissions and
  14# limitations under the License.
  15
  16import codecs
  17import io
  18import os
  19import tempfile
  20from collections.abc import Callable, Iterator
  21from typing import IO, Any, TypeVar
  22
  23import boto3
  24import botocore
  25from boto3.exceptions import S3UploadFailedError
  26from boto3.s3.transfer import TransferConfig
  27from botocore.credentials import RefreshableCredentials
  28from botocore.exceptions import ClientError, IncompleteReadError, ReadTimeoutError, ResponseStreamingError
  29from botocore.session import get_session
  30from dateutil.parser import parse as dateutil_parse
  31
  32from multistorageclient_rust import RustClient, RustClientError, RustRetryableError
  33
  34from ..constants import DEFAULT_CONNECT_TIMEOUT, DEFAULT_MAX_POOL_CONNECTIONS, DEFAULT_READ_TIMEOUT
  35from ..rust_utils import parse_retry_config, run_async_rust_client_method
  36from ..signers import CloudFrontURLSigner, URLSigner
  37from ..telemetry import Telemetry
  38from ..types import (
  39    AWARE_DATETIME_MIN,
  40    Credentials,
  41    CredentialsProvider,
  42    ObjectMetadata,
  43    PreconditionFailedError,
  44    Range,
  45    RetryableError,
  46    SignerType,
  47    SymlinkHandling,
  48)
  49from ..utils import (
  50    safe_makedirs,
  51    split_path,
  52    validate_attributes,
  53)
  54from .base import BaseStorageProvider
  55
  56_T = TypeVar("_T")
  57
  58MiB = 1024 * 1024
  59
  60# Python and Rust share the same multipart_threshold to keep the code simple.
  61MULTIPART_THRESHOLD = 64 * MiB
  62MULTIPART_CHUNKSIZE = 32 * MiB
  63IO_CHUNKSIZE = 32 * MiB
  64PYTHON_MAX_CONCURRENCY = 8
  65MAX_POOL_CONNECTIONS = DEFAULT_MAX_POOL_CONNECTIONS
  66
  67PROVIDER = "s3"
  68
  69EXPRESS_ONEZONE_STORAGE_CLASS = "EXPRESS_ONEZONE"
  70
  71# Source: https://docs.aws.amazon.com/AmazonS3/latest/userguide/checking-object-integrity.html
  72SUPPORTED_CHECKSUM_ALGORITHMS = frozenset({"CRC32", "CRC32C", "SHA1", "SHA256", "CRC64NVME"})
  73
  74
[docs] 75class StaticS3CredentialsProvider(CredentialsProvider): 76 """ 77 A concrete implementation of the :py:class:`multistorageclient.types.CredentialsProvider` that provides static S3 credentials. 78 """ 79 80 _access_key: str 81 _secret_key: str 82 _session_token: str | None 83 84 def __init__(self, access_key: str, secret_key: str, session_token: str | None = None): 85 """ 86 Initializes the :py:class:`StaticS3CredentialsProvider` with the provided access key, secret key, and optional 87 session token. 88 89 :param access_key: The access key for S3 authentication. 90 :param secret_key: The secret key for S3 authentication. 91 :param session_token: An optional session token for temporary credentials. 92 """ 93 self._access_key = access_key 94 self._secret_key = secret_key 95 self._session_token = session_token 96
[docs] 97 def get_credentials(self) -> Credentials: 98 return Credentials( 99 access_key=self._access_key, 100 secret_key=self._secret_key, 101 token=self._session_token, 102 expiration=None, 103 )
104
[docs] 105 def refresh_credentials(self) -> None: 106 pass
107 108 109DEFAULT_PRESIGN_EXPIRES_IN = 3600 110 111_S3_METHOD_MAPPING: dict[str, str] = { 112 "GET": "get_object", 113 "PUT": "put_object", 114} 115 116
[docs] 117class S3URLSigner(URLSigner): 118 """Generates pre-signed URLs using the boto3 S3 client. 119 120 When the underlying credentials are temporary (STS, IAM role, EC2 instance 121 profile), the effective URL lifetime is the **shorter** of ``expires_in`` 122 and the remaining credential lifetime — boto3 will not warn if the 123 credential expires before ``expires_in``. 124 125 See https://docs.aws.amazon.com/AmazonS3/latest/userguide/using-presigned-url.html 126 """ 127 128 def __init__(self, s3_client: Any, bucket: str, expires_in: int = DEFAULT_PRESIGN_EXPIRES_IN) -> None: 129 self._s3_client = s3_client 130 self._bucket = bucket 131 self._expires_in = expires_in 132
[docs] 133 def generate_presigned_url(self, path: str, *, method: str = "GET") -> str: 134 client_method = _S3_METHOD_MAPPING.get(method.upper()) 135 if client_method is None: 136 raise ValueError(f"Unsupported method for S3 presigning: {method!r}") 137 return self._s3_client.generate_presigned_url( 138 ClientMethod=client_method, 139 Params={"Bucket": self._bucket, "Key": path}, 140 ExpiresIn=self._expires_in, 141 )
142 143
[docs] 144class S3StorageProvider(BaseStorageProvider): 145 """ 146 A concrete implementation of the :py:class:`multistorageclient.types.StorageProvider` for interacting with Amazon S3 or S3-compatible object stores. 147 """ 148 149 def __init__( 150 self, 151 region_name: str = "", 152 endpoint_url: str = "", 153 base_path: str = "", 154 credentials_provider: CredentialsProvider | None = None, 155 profile_name: str | None = None, 156 config_dict: dict[str, Any] | None = None, 157 telemetry_provider: Callable[[], Telemetry] | None = None, 158 verify: bool | str | None = None, 159 **kwargs: Any, 160 ) -> None: 161 """ 162 Initializes the :py:class:`S3StorageProvider` with the region, endpoint URL, and optional credentials provider. 163 164 :param region_name: The AWS region where the S3 bucket is located. 165 :param endpoint_url: The custom endpoint URL for the S3 service. 166 :param base_path: The root prefix path within the S3 bucket where all operations will be scoped. 167 :param credentials_provider: The provider to retrieve S3 credentials. 168 :param profile_name: AWS shared configuration + credentials files profile. For :py:class:`boto3.session.Session`. Ignored if ``credentials_provider`` is set. 169 :param config_dict: Resolved MSC config. 170 :param telemetry_provider: A function that provides a telemetry instance. 171 :param verify: Controls SSL certificate verification. 172 Can be ``True`` (verify using system CA bundle, default), ``False`` (skip verification), or a string path to a custom CA certificate bundle. 173 :param request_checksum_calculation: For :py:class:`botocore.config.Config`. 174 When the underlying S3 client should calculate request checksums. 175 :param response_checksum_validation: For :py:class:`botocore.config.Config`. 176 When the underlying S3 client should validate response checksums. 177 :param max_pool_connections: For :py:class:`botocore.config.Config`. 178 The maximum number of connections to keep in a connection pool. 179 :param connect_timeout: For :py:class:`botocore.config.Config`. 180 The time in seconds till a timeout exception is thrown when attempting to make a connection. 181 :param read_timeout: For :py:class:`botocore.config.Config`. 182 The time in seconds till a timeout exception is thrown when attempting to read from a connection. 183 :param retries: For :py:class:`botocore.config.Config`. 184 A dictionary for configuration related to retry behavior. 185 :param s3: For :py:class:`botocore.config.Config`. 186 A dictionary of S3 specific configurations. 187 :param checksum_algorithm: Upload-only object integrity algorithm. 188 One of ``"CRC32"``, ``"CRC32C"``, ``"SHA1"``, ``"SHA256"``, ``"CRC64NVME"`` (case-insensitive). 189 When ``rust_client`` is enabled, only ``"SHA256"`` is accepted. 190 :param multipart_threshold: For :py:class:`boto3.s3.transfer.TransferConfig`. 191 The transfer size threshold for which multipart uploads, downloads, and copies will automatically be triggered. 192 :param max_concurrency: For :py:class:`boto3.s3.transfer.TransferConfig`. 193 The maximum number of threads that will be making requests to perform a transfer. 194 If ``use_threads`` is set to ``False``, the value provided is ignored as the transfer will only ever use the current thread. 195 :param multipart_chunksize: For :py:class:`boto3.s3.transfer.TransferConfig`. 196 The partition size of each part for a multipart transfer. 197 :param io_chunksize: For :py:class:`boto3.s3.transfer.TransferConfig`. 198 The max size of each chunk in the ``io`` queue. Currently, this is the size used when read is called on the downloaded stream as well. 199 Note: This value is ignored when resolved transfer manager type is CRTTransferManager. 200 """ 201 super().__init__( 202 base_path=base_path, 203 provider_name=PROVIDER, 204 config_dict=config_dict, 205 telemetry_provider=telemetry_provider, 206 ) 207 208 self._region_name = region_name 209 self._endpoint_url = endpoint_url 210 self._credentials_provider = credentials_provider 211 self._profile_name = profile_name 212 self._verify = verify 213 214 self._signature_version = kwargs.get("signature_version", "s3v4") 215 self._s3_client = self._create_s3_client( 216 request_checksum_calculation=kwargs.get("request_checksum_calculation"), 217 response_checksum_validation=kwargs.get("response_checksum_validation"), 218 max_pool_connections=kwargs.get("max_pool_connections", MAX_POOL_CONNECTIONS), 219 connect_timeout=kwargs.get("connect_timeout", DEFAULT_CONNECT_TIMEOUT), 220 read_timeout=kwargs.get("read_timeout", DEFAULT_READ_TIMEOUT), 221 retries=kwargs.get("retries"), 222 s3=kwargs.get("s3"), 223 ) 224 self._transfer_config = TransferConfig( 225 multipart_threshold=int(kwargs.get("multipart_threshold", MULTIPART_THRESHOLD)), 226 max_concurrency=int(kwargs.get("max_concurrency", PYTHON_MAX_CONCURRENCY)), 227 multipart_chunksize=int(kwargs.get("multipart_chunksize", MULTIPART_CHUNKSIZE)), 228 io_chunksize=int(kwargs.get("io_chunksize", IO_CHUNKSIZE)), 229 use_threads=True, 230 ) 231 self._multipart_threshold = self._transfer_config.multipart_threshold 232 233 self._signer_cache: dict[tuple, URLSigner] = {} 234 235 self._checksum_algorithm: str | None = self._validate_checksum_algorithm(kwargs.get("checksum_algorithm")) 236 237 self._rust_client = None 238 if "rust_client" in kwargs: 239 # Inherit the rust client options from the kwargs 240 rust_client_options = kwargs["rust_client"] 241 if "max_pool_connections" in kwargs: 242 rust_client_options["max_pool_connections"] = kwargs["max_pool_connections"] 243 if "max_concurrency" in kwargs: 244 rust_client_options["max_concurrency"] = kwargs["max_concurrency"] 245 if "multipart_chunksize" in kwargs: 246 rust_client_options["multipart_chunksize"] = kwargs["multipart_chunksize"] 247 if "read_timeout" in kwargs: 248 rust_client_options["read_timeout"] = kwargs["read_timeout"] 249 if "connect_timeout" in kwargs: 250 rust_client_options["connect_timeout"] = kwargs["connect_timeout"] 251 if "checksum_algorithm" in kwargs: 252 rust_client_options["checksum_algorithm"] = kwargs["checksum_algorithm"] 253 if self._signature_version == "UNSIGNED": 254 rust_client_options["skip_signature"] = True 255 256 # Rust client only supports SHA256 today (object_store limitation). 257 rust_checksum = rust_client_options.get("checksum_algorithm") 258 if rust_checksum is not None: 259 if str(rust_checksum).upper() != "SHA256": 260 raise ValueError( 261 f"Rust client only supports checksum_algorithm='SHA256' today " 262 f"(object_store limitation), got '{rust_checksum!r}'. " 263 f"Disable rust_client or use SHA256." 264 ) 265 rust_client_options["checksum_algorithm"] = "sha256" 266 self._rust_client = self._create_rust_client(rust_client_options) 267 268 @staticmethod 269 def _validate_checksum_algorithm(value: str | None) -> str | None: 270 """ 271 Normalizes the optional ``checksum_algorithm`` to uppercase, or returns ``None`` if disabled. 272 """ 273 if value is None: 274 return None 275 if not isinstance(value, str) or not value: 276 raise ValueError( 277 f"checksum_algorithm must be one of {sorted(SUPPORTED_CHECKSUM_ALGORITHMS)} or None, got {value!r}." 278 ) 279 normalized = value.upper() 280 if normalized not in SUPPORTED_CHECKSUM_ALGORITHMS: 281 raise ValueError( 282 f"Unsupported checksum_algorithm '{value}'. Supported: {sorted(SUPPORTED_CHECKSUM_ALGORITHMS)}." 283 ) 284 return normalized 285 286 def _is_directory_bucket(self, bucket: str) -> bool: 287 """ 288 Determines if the bucket is a directory bucket based on bucket name. 289 """ 290 # S3 Express buckets have a specific naming convention 291 return "--x-s3" in bucket 292 293 def _create_s3_client( 294 self, 295 request_checksum_calculation: str | None = None, 296 response_checksum_validation: str | None = None, 297 max_pool_connections: int = MAX_POOL_CONNECTIONS, 298 connect_timeout: float | None = None, 299 read_timeout: float | None = None, 300 retries: dict[str, Any] | None = None, 301 s3: dict[str, Any] | None = None, 302 ): 303 """ 304 Creates and configures the boto3 S3 client, using refreshable credentials if possible. 305 306 :param request_checksum_calculation: For :py:class:`botocore.config.Config`. When the underlying S3 client should calculate request checksums. See the equivalent option in the `AWS configuration file <https://boto3.amazonaws.com/v1/documentation/api/latest/guide/configuration.html#using-a-configuration-file>`_. 307 :param response_checksum_validation: For :py:class:`botocore.config.Config`. When the underlying S3 client should validate response checksums. See the equivalent option in the `AWS configuration file <https://boto3.amazonaws.com/v1/documentation/api/latest/guide/configuration.html#using-a-configuration-file>`_. 308 :param max_pool_connections: For :py:class:`botocore.config.Config`. The maximum number of connections to keep in a connection pool. 309 :param connect_timeout: For :py:class:`botocore.config.Config`. The time in seconds till a timeout exception is thrown when attempting to make a connection. 310 :param read_timeout: For :py:class:`botocore.config.Config`. The time in seconds till a timeout exception is thrown when attempting to read from a connection. 311 :param retries: For :py:class:`botocore.config.Config`. A dictionary for configuration related to retry behavior. 312 :param s3: For :py:class:`botocore.config.Config`. A dictionary of S3 specific configurations. 313 314 :return: The configured S3 client. 315 """ 316 options = { 317 # https://botocore.amazonaws.com/v1/documentation/api/latest/reference/config.html 318 "config": boto3.session.Config( # pyright: ignore [reportAttributeAccessIssue] 319 max_pool_connections=max_pool_connections, 320 connect_timeout=connect_timeout, 321 read_timeout=read_timeout, 322 retries=retries or {"mode": "standard"}, 323 request_checksum_calculation=request_checksum_calculation, 324 response_checksum_validation=response_checksum_validation, 325 s3=s3, 326 ), 327 } 328 329 if self._region_name: 330 options["region_name"] = self._region_name 331 332 if self._endpoint_url: 333 options["endpoint_url"] = self._endpoint_url 334 335 if self._verify is not None: 336 options["verify"] = self._verify 337 338 if self._credentials_provider: 339 creds = self._fetch_credentials() 340 if creds.get("expiry_time"): 341 # Use RefreshableCredentials if expiry_time provided. 342 refreshable_credentials = RefreshableCredentials.create_from_metadata( 343 metadata=creds, refresh_using=self._fetch_credentials, method="custom-refresh" 344 ) 345 346 botocore_session = get_session() 347 botocore_session._credentials = refreshable_credentials 348 349 boto3_session = boto3.Session(botocore_session=botocore_session) 350 351 return boto3_session.client("s3", **options) 352 else: 353 # Add static credentials to the options dictionary 354 options["aws_access_key_id"] = creds["access_key"] 355 options["aws_secret_access_key"] = creds["secret_key"] 356 if creds["token"]: 357 options["aws_session_token"] = creds["token"] 358 359 if self._signature_version: 360 signature_config = botocore.config.Config( # pyright: ignore[reportAttributeAccessIssue] 361 signature_version=botocore.UNSIGNED 362 if self._signature_version == "UNSIGNED" 363 else self._signature_version 364 ) 365 options["config"] = options["config"].merge(signature_config) 366 367 # Fallback to standard credential chain. 368 session = boto3.Session(profile_name=self._profile_name) 369 return session.client("s3", **options) 370 371 def _create_rust_client(self, rust_client_options: dict[str, Any] | None = None): 372 """ 373 Creates and configures the rust client, using refreshable credentials if possible. 374 """ 375 configs = dict(rust_client_options) if rust_client_options else {} 376 377 # Extract and parse retry configuration 378 retry_config = parse_retry_config(configs) 379 380 if self._region_name and "region_name" not in configs: 381 configs["region_name"] = self._region_name 382 383 if self._endpoint_url and "endpoint_url" not in configs: 384 configs["endpoint_url"] = self._endpoint_url 385 386 if self._profile_name and "profile_name" not in configs: 387 configs["profile_name"] = self._profile_name 388 389 # If the user specifies a bucket, use it. Otherwise, use the base path. 390 if "bucket" not in configs: 391 bucket, _ = split_path(self._base_path) 392 configs["bucket"] = bucket 393 394 if "max_pool_connections" not in configs: 395 configs["max_pool_connections"] = MAX_POOL_CONNECTIONS 396 397 return RustClient( 398 provider=PROVIDER, 399 configs=configs, 400 credentials_provider=self._credentials_provider, 401 retry=retry_config, 402 ) 403 404 def _fetch_credentials(self) -> dict: 405 """ 406 Refreshes the S3 client if the current credentials are expired. 407 """ 408 if not self._credentials_provider: 409 raise RuntimeError("Cannot fetch credentials if no credential provider configured.") 410 self._credentials_provider.refresh_credentials() 411 credentials = self._credentials_provider.get_credentials() 412 return { 413 "access_key": credentials.access_key, 414 "secret_key": credentials.secret_key, 415 "token": credentials.token, 416 "expiry_time": credentials.expiration, 417 } 418 419 def _translate_errors( 420 self, 421 func: Callable[[], _T], 422 operation: str, 423 bucket: str, 424 key: str, 425 ) -> _T: 426 """ 427 Translates errors like timeouts and client errors. 428 429 :param func: The function that performs the actual S3 operation. 430 :param operation: The type of operation being performed (e.g., "PUT", "GET", "DELETE"). 431 :param bucket: The name of the S3 bucket involved in the operation. 432 :param key: The key of the object within the S3 bucket. 433 434 :return: The result of the S3 operation, typically the return value of the `func` callable. 435 """ 436 try: 437 return func() 438 except ClientError as error: 439 status_code = error.response["ResponseMetadata"]["HTTPStatusCode"] 440 request_id = error.response["ResponseMetadata"].get("RequestId") 441 host_id = error.response["ResponseMetadata"].get("HostId") 442 error_code = error.response["Error"]["Code"] 443 error_info = f"request_id: {request_id}, host_id: {host_id}, status_code: {status_code}" 444 445 if status_code == 404: 446 if error_code == "NoSuchUpload": 447 error_message = error.response["Error"]["Message"] 448 raise RetryableError(f"Multipart upload failed for {bucket}/{key}: {error_message}") from error 449 raise FileNotFoundError(f"Object {bucket}/{key} does not exist. {error_info}") # pylint: disable=raise-missing-from 450 elif status_code == 403: 451 raise PermissionError( 452 f"Permission denied to {operation} object(s) at {bucket}/{key}. {error_info}" 453 ) from error 454 elif status_code == 412: # Precondition Failed 455 raise PreconditionFailedError( 456 f"ETag mismatch for {operation} operation on {bucket}/{key}. {error_info}" 457 ) from error 458 elif status_code == 429: 459 raise RetryableError( 460 f"Too many request to {operation} object(s) at {bucket}/{key}. {error_info}" 461 ) from error 462 elif status_code == 503: 463 raise RetryableError( 464 f"Service unavailable when {operation} object(s) at {bucket}/{key}. {error_info}" 465 ) from error 466 elif status_code == 501: 467 raise NotImplementedError( 468 f"Operation {operation} not implemented for object(s) at {bucket}/{key}. {error_info}" 469 ) from error 470 elif status_code == 408: 471 # 408 Request Timeout is from Google Cloud Storage 472 raise RetryableError( 473 f"Request timeout when {operation} object(s) at {bucket}/{key}. {error_info}" 474 ) from error 475 else: 476 raise RuntimeError( 477 f"Failed to {operation} object(s) at {bucket}/{key}. {error_info}, " 478 f"error_type: {type(error).__name__}" 479 ) from error 480 except RustClientError as error: 481 message = error.args[0] 482 status_code = error.args[1] 483 if status_code == 404: 484 raise FileNotFoundError(f"Object {bucket}/{key} does not exist. {message}") from error 485 elif status_code == 403: 486 raise PermissionError( 487 f"Permission denied to {operation} object(s) at {bucket}/{key}. {message}" 488 ) from error 489 else: 490 raise RetryableError( 491 f"Failed to {operation} object(s) at {bucket}/{key}. {message}. status_code: {status_code}" 492 ) from error 493 except (ReadTimeoutError, IncompleteReadError, ResponseStreamingError) as error: 494 raise RetryableError( 495 f"Failed to {operation} object(s) at {bucket}/{key} due to network timeout or incomplete read. " 496 f"error_type: {type(error).__name__}" 497 ) from error 498 except RustRetryableError as error: 499 raise RetryableError( 500 f"Failed to {operation} object(s) at {bucket}/{key} due to retryable error from Rust. " 501 f"error_type: {type(error).__name__}" 502 ) from error 503 except Exception as error: 504 if ( 505 isinstance(error, S3UploadFailedError) 506 and isinstance(error.__context__, ClientError) 507 and error.__context__.response["ResponseMetadata"]["HTTPStatusCode"] == 403 508 ): 509 raise PermissionError( 510 f"Permission denied to {operation} object(s) at {bucket}/{key}. {error}" 511 ) from error 512 raise RuntimeError( 513 f"Failed to {operation} object(s) at {bucket}/{key}, error type: {type(error).__name__}, error: {error}" 514 ) from error 515 516 def _put_object( 517 self, 518 path: str, 519 body: bytes, 520 if_match: str | None = None, 521 if_none_match: str | None = None, 522 attributes: dict[str, str] | None = None, 523 content_type: str | None = None, 524 ) -> int: 525 """ 526 Uploads an object to the specified S3 path. 527 528 :param path: The S3 path where the object will be uploaded. 529 :param body: The content of the object as bytes. 530 :param if_match: Optional If-Match header value. Use "*" to only upload if the object doesn't exist. 531 :param if_none_match: Optional If-None-Match header value. Use "*" to only upload if the object doesn't exist. 532 :param attributes: Optional attributes to attach to the object. 533 :param content_type: Optional Content-Type header value. 534 """ 535 bucket, key = split_path(path) 536 537 def _invoke_api() -> int: 538 kwargs = {"Bucket": bucket, "Key": key, "Body": body} 539 if content_type: 540 kwargs["ContentType"] = content_type 541 if self._is_directory_bucket(bucket): 542 kwargs["StorageClass"] = EXPRESS_ONEZONE_STORAGE_CLASS 543 if if_match: 544 kwargs["IfMatch"] = if_match 545 if if_none_match: 546 kwargs["IfNoneMatch"] = if_none_match 547 validated_attributes = validate_attributes(attributes) 548 if validated_attributes: 549 kwargs["Metadata"] = validated_attributes 550 551 # TODO(NGCDP-5804): Add support to update ContentType header in Rust client 552 rust_unsupported_feature_keys = {"StorageClass", "IfMatch", "IfNoneMatch", "ContentType"} 553 if ( 554 self._rust_client 555 # Rust client doesn't support creating objects with trailing /, see https://github.com/apache/arrow-rs/issues/7026 556 and not path.endswith("/") 557 and all(key not in kwargs for key in rust_unsupported_feature_keys) 558 ): 559 rust_attributes = {"attributes": validated_attributes} if validated_attributes else {} 560 run_async_rust_client_method(self._rust_client, "put", key, body, **rust_attributes) 561 else: 562 if self._checksum_algorithm: 563 kwargs["ChecksumAlgorithm"] = self._checksum_algorithm 564 self._s3_client.put_object(**kwargs) 565 566 return len(body) 567 568 return self._translate_errors(_invoke_api, operation="PUT", bucket=bucket, key=key) 569 570 def _get_object(self, path: str, byte_range: Range | None = None) -> bytes: 571 bucket, key = split_path(path) 572 573 def _invoke_api() -> bytes: 574 if byte_range: 575 bytes_range = f"bytes={byte_range.offset}-{byte_range.offset + byte_range.size - 1}" 576 if self._rust_client: 577 response = run_async_rust_client_method( 578 self._rust_client, 579 "get", 580 key, 581 byte_range, 582 ) 583 return response 584 else: 585 response = self._s3_client.get_object(Bucket=bucket, Key=key, Range=bytes_range) 586 else: 587 if self._rust_client: 588 response = run_async_rust_client_method(self._rust_client, "get", key) 589 return response 590 else: 591 response = self._s3_client.get_object(Bucket=bucket, Key=key) 592 593 return response["Body"].read() 594 595 return self._translate_errors(_invoke_api, operation="GET", bucket=bucket, key=key) 596 597 def _copy_object(self, src_path: str, dest_path: str) -> int: 598 src_bucket, src_key = split_path(src_path) 599 dest_bucket, dest_key = split_path(dest_path) 600 601 src_object = self._get_object_metadata(src_path) 602 603 def _invoke_api() -> int: 604 self._s3_client.copy( 605 CopySource={"Bucket": src_bucket, "Key": src_key}, 606 Bucket=dest_bucket, 607 Key=dest_key, 608 Config=self._transfer_config, 609 ) 610 611 return src_object.content_length 612 613 return self._translate_errors(_invoke_api, operation="COPY", bucket=dest_bucket, key=dest_key) 614 615 def _delete_object(self, path: str, if_match: str | None = None) -> None: 616 bucket, key = split_path(path) 617 618 def _invoke_api() -> None: 619 # Delete conditionally when if_match (etag) is provided; otherwise delete unconditionally 620 if if_match: 621 self._s3_client.delete_object(Bucket=bucket, Key=key, IfMatch=if_match) 622 else: 623 self._s3_client.delete_object(Bucket=bucket, Key=key) 624 625 return self._translate_errors(_invoke_api, operation="DELETE", bucket=bucket, key=key) 626 627 def _delete_objects(self, paths: list[str]) -> None: 628 if not paths: 629 return 630 631 by_bucket: dict[str, list[str]] = {} 632 for p in paths: 633 bucket, key = split_path(p) 634 by_bucket.setdefault(bucket, []).append(key) 635 636 S3_BATCH_LIMIT = 1000 637 638 def _invoke_api() -> None: 639 all_errors: list[str] = [] 640 for bucket, keys in by_bucket.items(): 641 for i in range(0, len(keys), S3_BATCH_LIMIT): 642 chunk = keys[i : i + S3_BATCH_LIMIT] 643 response = self._s3_client.delete_objects( 644 Bucket=bucket, Delete={"Objects": [{"Key": k} for k in chunk]} 645 ) 646 errors = response.get("Errors") or [] 647 for e in errors: 648 all_errors.append(f"{bucket}/{e.get('Key', '?')}: {e.get('Code', '')} {e.get('Message', '')}") 649 if all_errors: 650 raise RuntimeError(f"DeleteObjects reported errors: {'; '.join(all_errors)}") 651 652 bucket_desc = "(" + "|".join(by_bucket) + ")" 653 key_desc = "(" + "|".join(str(len(keys)) for keys in by_bucket.values()) + " keys)" 654 self._translate_errors(_invoke_api, operation="DELETE_MANY", bucket=bucket_desc, key=key_desc) 655 656 def _is_dir(self, path: str) -> bool: 657 # Ensure the path ends with '/' to mimic a directory 658 path = self._append_delimiter(path) 659 660 bucket, key = split_path(path) 661 662 def _invoke_api() -> bool: 663 # List objects with the given prefix 664 response = self._s3_client.list_objects_v2(Bucket=bucket, Prefix=key, MaxKeys=1, Delimiter="/") 665 666 # Check if there are any contents or common prefixes 667 return bool(response.get("Contents", []) or response.get("CommonPrefixes", [])) 668 669 return self._translate_errors(_invoke_api, operation="LIST", bucket=bucket, key=key) 670 671 def _make_symlink(self, path: str, target: str) -> None: 672 bucket, key = split_path(path) 673 target_bucket, target_key = split_path(target) 674 if bucket != target_bucket: 675 raise ValueError(f"Cannot create cross-bucket symlink: '{bucket}' -> '{target_bucket}'.") 676 relative_target = ObjectMetadata.encode_symlink_target(key, target_key) 677 678 def _invoke_api() -> None: 679 self._s3_client.put_object( 680 Bucket=bucket, 681 Key=key, 682 Body=b"", 683 Metadata={"msc-symlink-target": relative_target}, 684 ) 685 686 self._translate_errors(_invoke_api, operation="PUT", bucket=bucket, key=key) 687 688 def _get_object_metadata(self, path: str, strict: bool = True) -> ObjectMetadata: 689 bucket, key = split_path(path) 690 if path.endswith("/") or (bucket and not key): 691 # If path ends with "/" or empty key name is provided, then assume it's a "directory", 692 # which metadata is not guaranteed to exist for cases such as 693 # "virtual prefix" that was never explicitly created. 694 if self._is_dir(path): 695 return ObjectMetadata( 696 key=path, 697 type="directory", 698 content_length=0, 699 last_modified=AWARE_DATETIME_MIN, 700 ) 701 else: 702 raise FileNotFoundError(f"Directory {path} does not exist.") 703 else: 704 705 def _invoke_api() -> ObjectMetadata: 706 response = self._s3_client.head_object(Bucket=bucket, Key=key) 707 user_metadata = response.get("Metadata") 708 symlink_target = user_metadata.get("msc-symlink-target") if user_metadata else None 709 710 return ObjectMetadata( 711 key=path, 712 type="file", 713 content_length=response["ContentLength"], 714 content_type=response.get("ContentType"), 715 last_modified=response.get("LastModified", AWARE_DATETIME_MIN), 716 etag=response["ETag"].strip('"') if "ETag" in response else None, 717 storage_class=response.get("StorageClass"), 718 metadata=user_metadata, 719 symlink_target=symlink_target, 720 ) 721 722 try: 723 return self._translate_errors(_invoke_api, operation="HEAD", bucket=bucket, key=key) 724 except FileNotFoundError: 725 if strict: 726 # If the object does not exist on the given path, we will append a trailing slash and 727 # check if the path is a directory. 728 path = self._append_delimiter(path) 729 if self._is_dir(path): 730 return ObjectMetadata( 731 key=path, 732 type="directory", 733 content_length=0, 734 last_modified=AWARE_DATETIME_MIN, 735 ) 736 raise 737 738 def _list_objects( 739 self, 740 path: str, 741 start_after: str | None = None, 742 end_at: str | None = None, 743 include_directories: bool = False, 744 symlink_handling: SymlinkHandling = SymlinkHandling.FOLLOW, 745 ) -> Iterator[ObjectMetadata]: 746 bucket, prefix = split_path(path) 747 748 # Get the prefix of the start_after and end_at paths relative to the bucket. 749 if start_after: 750 _, start_after = split_path(start_after) 751 if end_at: 752 _, end_at = split_path(end_at) 753 754 def _invoke_api() -> Iterator[ObjectMetadata]: 755 paginator = self._s3_client.get_paginator("list_objects_v2") 756 if include_directories: 757 page_iterator = paginator.paginate( 758 Bucket=bucket, Prefix=prefix, Delimiter="/", StartAfter=(start_after or "") 759 ) 760 else: 761 page_iterator = paginator.paginate(Bucket=bucket, Prefix=prefix, StartAfter=(start_after or "")) 762 763 for page in page_iterator: 764 # A page holds both CommonPrefixes and Contents for the same key window, so a prefix past 765 # end_at must not stop the listing before the page's objects have been processed. 766 past_end_at = False 767 for item in page.get("CommonPrefixes", []): 768 prefix_key = item["Prefix"].rstrip("/") 769 # Filter by start_after and end_at - S3's StartAfter doesn't filter CommonPrefixes 770 if (start_after is None or start_after < prefix_key) and (end_at is None or prefix_key <= end_at): 771 yield ObjectMetadata( 772 key=os.path.join(bucket, prefix_key), 773 type="directory", 774 content_length=0, 775 last_modified=AWARE_DATETIME_MIN, 776 ) 777 elif end_at is not None and end_at < prefix_key: 778 past_end_at = True 779 break 780 781 # S3 guarantees lexicographical order for general purpose buckets (for 782 # normal S3) but not directory buckets (for S3 Express One Zone). 783 for response_object in page.get("Contents", []): 784 key = response_object["Key"] 785 if end_at is None or key <= end_at: 786 if key.endswith("/"): 787 if include_directories: 788 yield ObjectMetadata( 789 key=os.path.join(bucket, key.rstrip("/")), 790 type="directory", 791 content_length=0, 792 last_modified=response_object["LastModified"], 793 ) 794 else: 795 symlink_target = None 796 if response_object.get("Size", 0) == 0: 797 try: 798 meta = self._get_object_metadata(os.path.join(bucket, key)) 799 symlink_target = meta.symlink_target 800 except Exception: 801 symlink_target = None 802 yield ObjectMetadata( 803 key=os.path.join(bucket, key), 804 type="file", 805 content_length=response_object["Size"], 806 last_modified=response_object["LastModified"], 807 etag=response_object["ETag"].strip('"'), 808 storage_class=response_object.get("StorageClass"), 809 symlink_target=symlink_target, 810 ) 811 else: 812 return 813 814 if past_end_at: 815 return 816 817 return self._translate_errors(_invoke_api, operation="LIST", bucket=bucket, key=prefix) 818 819 @property 820 def supports_parallel_listing(self) -> bool: 821 """ 822 S3 supports parallel listing via delimiter-based prefix discovery. 823 824 Note: Directory bucket handling is done in list_objects_recursive(). 825 """ 826 return True 827
[docs] 828 def list_objects_recursive( 829 self, 830 path: str = "", 831 start_after: str | None = None, 832 end_at: str | None = None, 833 max_workers: int = 32, 834 look_ahead: int = 2, 835 symlink_handling: SymlinkHandling = SymlinkHandling.FOLLOW, 836 ) -> Iterator[ObjectMetadata]: 837 """ 838 List all objects recursively using parallel prefix discovery for improved performance. 839 840 For S3, uses the Rust client's list_recursive when available for maximum performance. 841 Falls back to Python implementation otherwise. 842 843 Returns files only (no directories), in lexicographic order. 844 845 :param symlink_handling: How to handle symbolic links during listing. 846 """ 847 if (start_after is not None) and (end_at is not None) and not (start_after < end_at): 848 raise ValueError(f"start_after ({start_after}) must be before end_at ({end_at})!") 849 850 full_path = self._prepend_base_path(path) 851 bucket, _ = split_path(full_path) 852 853 if self._is_directory_bucket(bucket): 854 yield from self.list_objects( 855 path, start_after, end_at, include_directories=False, symlink_handling=symlink_handling 856 ) 857 return 858 859 if self._rust_client: 860 yield from self._emit_metrics( 861 operation=BaseStorageProvider._Operation.LIST, 862 f=lambda: self._list_objects_recursive_rust(path, full_path, bucket, start_after, end_at, max_workers), 863 ) 864 else: 865 yield from super().list_objects_recursive( 866 path, start_after, end_at, max_workers, look_ahead, symlink_handling 867 )
868 869 def _list_objects_recursive_rust( 870 self, 871 path: str, 872 full_path: str, 873 bucket: str, 874 start_after: str | None, 875 end_at: str | None, 876 max_workers: int, 877 ) -> Iterator[ObjectMetadata]: 878 """ 879 Use Rust client's list_recursive for parallel listing. 880 881 The Rust client already handles parallel listing internally. 882 Returns files only in lexicographic order. 883 """ 884 _, prefix = split_path(full_path) 885 886 def _invoke_api() -> Iterator[ObjectMetadata]: 887 result = run_async_rust_client_method( 888 self._rust_client, 889 "list_recursive", 890 [prefix] if prefix else [""], 891 max_concurrency=max_workers, 892 ) 893 894 start_after_full = self._prepend_base_path(start_after) if start_after else None 895 end_at_full = self._prepend_base_path(end_at) if end_at else None 896 897 for obj in result.objects: 898 full_key = os.path.join(bucket, obj.key) 899 900 if start_after_full and full_key <= start_after_full: 901 continue 902 if end_at_full and full_key > end_at_full: 903 break 904 905 relative_key = full_key.removeprefix(self._base_path).lstrip("/") 906 907 yield ObjectMetadata( 908 key=relative_key, 909 content_length=obj.content_length, 910 last_modified=dateutil_parse(obj.last_modified), 911 type="file" if obj.object_type == "object" else obj.object_type, 912 etag=obj.etag, 913 ) 914 915 yield from self._translate_errors(_invoke_api, operation="LIST", bucket=bucket, key=prefix) 916 917 def _upload_file( 918 self, 919 remote_path: str, 920 f: str | IO, 921 attributes: dict[str, str] | None = None, 922 content_type: str | None = None, 923 ) -> int: 924 file_size: int = 0 925 926 if isinstance(f, str): 927 bucket, key = split_path(remote_path) 928 file_size = os.path.getsize(f) 929 930 # Upload small files 931 if file_size <= self._multipart_threshold: 932 if self._rust_client and not content_type and not self._is_directory_bucket(bucket): 933 validated_attributes = validate_attributes(attributes) 934 rust_attributes = {"attributes": validated_attributes} if validated_attributes else {} 935 self._translate_errors( 936 lambda: run_async_rust_client_method(self._rust_client, "upload", f, key, **rust_attributes), 937 operation="PUT", 938 bucket=bucket, 939 key=key, 940 ) 941 else: 942 with open(f, "rb") as fp: 943 self._put_object(remote_path, fp.read(), attributes=attributes, content_type=content_type) 944 return file_size 945 946 # Upload large files using TransferConfig 947 def _invoke_api() -> int: 948 extra_args = {} 949 if content_type: 950 extra_args["ContentType"] = content_type 951 if self._is_directory_bucket(bucket): 952 extra_args["StorageClass"] = EXPRESS_ONEZONE_STORAGE_CLASS 953 validated_attributes = validate_attributes(attributes) 954 if validated_attributes: 955 extra_args["Metadata"] = validated_attributes 956 if self._rust_client and "ContentType" not in extra_args and "StorageClass" not in extra_args: 957 rust_attributes = {"attributes": validated_attributes} if validated_attributes else {} 958 run_async_rust_client_method( 959 self._rust_client, "upload_multipart_from_file", f, key, **rust_attributes 960 ) 961 else: 962 if self._checksum_algorithm: 963 extra_args["ChecksumAlgorithm"] = self._checksum_algorithm 964 self._s3_client.upload_file( 965 Filename=f, 966 Bucket=bucket, 967 Key=key, 968 Config=self._transfer_config, 969 ExtraArgs=extra_args, 970 ) 971 972 return file_size 973 974 return self._translate_errors(_invoke_api, operation="PUT", bucket=bucket, key=key) 975 else: 976 # Upload small files 977 f.seek(0, io.SEEK_END) 978 file_size = f.tell() 979 f.seek(0) 980 981 if file_size <= self._multipart_threshold: 982 if isinstance(f, io.StringIO): 983 self._put_object( 984 remote_path, f.read().encode("utf-8"), attributes=attributes, content_type=content_type 985 ) 986 else: 987 self._put_object(remote_path, f.read(), attributes=attributes, content_type=content_type) 988 return file_size 989 990 # Upload large files using TransferConfig 991 bucket, key = split_path(remote_path) 992 993 def _invoke_api() -> int: 994 extra_args = {} 995 if content_type: 996 extra_args["ContentType"] = content_type 997 if self._is_directory_bucket(bucket): 998 extra_args["StorageClass"] = EXPRESS_ONEZONE_STORAGE_CLASS 999 validated_attributes = validate_attributes(attributes) 1000 if validated_attributes: 1001 extra_args["Metadata"] = validated_attributes 1002 1003 if ( 1004 self._rust_client 1005 and isinstance(f, io.BytesIO) 1006 and "ContentType" not in extra_args 1007 and "StorageClass" not in extra_args 1008 ): 1009 data = f.getbuffer() 1010 rust_attributes = {"attributes": validated_attributes} if validated_attributes else {} 1011 try: 1012 run_async_rust_client_method( 1013 self._rust_client, "upload_multipart_from_bytes", key, data, **rust_attributes 1014 ) 1015 finally: 1016 data.release() 1017 else: 1018 if self._checksum_algorithm: 1019 extra_args["ChecksumAlgorithm"] = self._checksum_algorithm 1020 self._s3_client.upload_fileobj( 1021 Fileobj=f, 1022 Bucket=bucket, 1023 Key=key, 1024 Config=self._transfer_config, 1025 ExtraArgs=extra_args, 1026 ) 1027 1028 return file_size 1029 1030 return self._translate_errors(_invoke_api, operation="PUT", bucket=bucket, key=key) 1031 1032 def _download_file(self, remote_path: str, f: str | IO, metadata: ObjectMetadata | None = None) -> int: 1033 if metadata is None: 1034 metadata = self._get_object_metadata(remote_path) 1035 1036 if isinstance(f, str): 1037 bucket, key = split_path(remote_path) 1038 if os.path.dirname(f): 1039 safe_makedirs(os.path.dirname(f)) 1040 1041 # Download small files 1042 if metadata.content_length <= self._multipart_threshold: 1043 if self._rust_client: 1044 run_async_rust_client_method(self._rust_client, "download", key, f) 1045 else: 1046 with tempfile.NamedTemporaryFile(mode="wb", delete=False, dir=os.path.dirname(f), prefix=".") as fp: 1047 temp_file_path = fp.name 1048 fp.write(self._get_object(remote_path)) 1049 os.rename(src=temp_file_path, dst=f) 1050 return metadata.content_length 1051 1052 # Download large files using TransferConfig 1053 def _invoke_api() -> int: 1054 with tempfile.NamedTemporaryFile(mode="wb", delete=False, dir=os.path.dirname(f), prefix=".") as fp: 1055 temp_file_path = fp.name 1056 if self._rust_client: 1057 run_async_rust_client_method( 1058 self._rust_client, "download_multipart_to_file", key, temp_file_path 1059 ) 1060 else: 1061 self._s3_client.download_fileobj( 1062 Bucket=bucket, 1063 Key=key, 1064 Fileobj=fp, 1065 Config=self._transfer_config, 1066 ) 1067 1068 os.rename(src=temp_file_path, dst=f) 1069 1070 return metadata.content_length 1071 1072 return self._translate_errors(_invoke_api, operation="GET", bucket=bucket, key=key) 1073 else: 1074 # Download small files 1075 if metadata.content_length <= self._multipart_threshold: 1076 response = self._get_object(remote_path) 1077 # Python client returns `bytes`, but Rust client returns an object that implements the buffer protocol, 1078 # so we need to check whether `.decode()` is available. 1079 if isinstance(f, io.StringIO): 1080 if hasattr(response, "decode"): 1081 f.write(response.decode("utf-8")) 1082 else: 1083 f.write(codecs.decode(memoryview(response), "utf-8")) 1084 else: 1085 f.write(response) 1086 return metadata.content_length 1087 1088 # Download large files using TransferConfig 1089 bucket, key = split_path(remote_path) 1090 1091 def _invoke_api() -> int: 1092 self._s3_client.download_fileobj( 1093 Bucket=bucket, 1094 Key=key, 1095 Fileobj=f, 1096 Config=self._transfer_config, 1097 ) 1098 1099 return metadata.content_length 1100 1101 return self._translate_errors(_invoke_api, operation="GET", bucket=bucket, key=key) 1102 1103 def _generate_presigned_url( 1104 self, 1105 path: str, 1106 *, 1107 method: str = "GET", 1108 signer_type: SignerType | None = None, 1109 signer_options: dict[str, Any] | None = None, 1110 ) -> str: 1111 options = signer_options or {} 1112 bucket, key = split_path(path) 1113 1114 if signer_type is None or signer_type == SignerType.S3: 1115 expires_in = int(options.get("expires_in", DEFAULT_PRESIGN_EXPIRES_IN)) 1116 cache_key: tuple = (SignerType.S3, bucket, expires_in) 1117 if cache_key not in self._signer_cache: 1118 self._signer_cache[cache_key] = S3URLSigner(self._s3_client, bucket, expires_in=expires_in) 1119 elif signer_type == SignerType.CLOUDFRONT: 1120 cache_key = (SignerType.CLOUDFRONT, frozenset(options.items())) 1121 if cache_key not in self._signer_cache: 1122 self._signer_cache[cache_key] = CloudFrontURLSigner(**options) 1123 else: 1124 raise ValueError(f"Unsupported signer type for S3 provider: {signer_type!r}") 1125 1126 return self._signer_cache[cache_key].generate_presigned_url(key, method=method)