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)