From 21dca66fbb63c18bc3f5255930f0abfb9be43ee6 Mon Sep 17 00:00:00 2001 From: xiaohongbo Date: Wed, 16 Sep 2026 02:16:43 -0700 Subject: [PATCH] [python] Prevent concurrent S3 metadata overwrites --- .github/workflows/ci-python.yml | 4 +- paimon-python/dev/lint-python.sh | 2 +- paimon-python/dev/requirements.txt | 1 + .../pypaimon/filesystem/pyarrow_file_io.py | 68 ++++++++ .../pypaimon/tests/s3_atomic_write_test.py | 163 ++++++++++++++++++ 5 files changed, 235 insertions(+), 3 deletions(-) create mode 100644 paimon-python/pypaimon/tests/s3_atomic_write_test.py diff --git a/.github/workflows/ci-python.yml b/.github/workflows/ci-python.yml index f7e127f07496..9d98a01c899d 100644 --- a/.github/workflows/ci-python.yml +++ b/.github/workflows/ci-python.yml @@ -122,7 +122,7 @@ jobs: if [[ "${{ matrix.python-version }}" == "3.6.15" ]]; then python -m pip install --upgrade pip==21.3.1 python --version - python -m pip install --no-cache-dir pyroaring readerwriterlock==1.0.9 'fsspec==2021.10.1' 'cachetools==4.2.4' 'ossfs==2021.8.0' pyarrow==6.0.1 pandas==1.1.5 'polars==0.9.12' 'fastavro==1.4.7' zstandard==0.19.0 dataclasses==0.8.0 flake8 pytest py4j==0.10.9.9 requests parameterized==0.8.1 datasketches==4.1.0 2>&1 >/dev/null + python -m pip install --no-cache-dir pyroaring readerwriterlock==1.0.9 'fsspec==2021.10.1' 'cachetools==4.2.4' 'ossfs==2021.8.0' pyarrow==6.0.1 pandas==1.1.5 'polars==0.9.12' 'fastavro==1.4.7' zstandard==0.19.0 dataclasses==0.8.0 boto3 flake8 pytest py4j==0.10.9.9 requests parameterized==0.8.1 datasketches==4.1.0 2>&1 >/dev/null python -m pip install 'lumina-data>=${{ env.LUMINA_DATA_VERSION }}' -i https://pypi.org/simple/ elif [[ "${{ matrix.python-version }}" == "3.7" ]]; then # 3.7 installs the version-pinned set declared for 3.7 in dev/requirements.txt. @@ -134,7 +134,7 @@ jobs: else python -m pip install --upgrade pip pip install torch --index-url https://download.pytorch.org/whl/cpu - python -m pip install pyroaring readerwriterlock==1.0.9 fsspec==2024.3.1 cachetools==5.3.3 ossfs==2023.12.0 ray==2.54.0 fastavro==1.11.1 'isal>=1.8,<2' zstandard==0.24.0 polars==1.32.0 duckdb==1.3.2 pylance==0.39.0 cramjam pytest~=7.0 py4j==0.10.9.9 requests parameterized==0.9.0 'daft>=0.7.6' 'datafusion>=54,<55' datasketches + python -m pip install pyroaring readerwriterlock==1.0.9 fsspec==2024.3.1 cachetools==5.3.3 ossfs==2023.12.0 ray==2.54.0 fastavro==1.11.1 'isal>=1.8,<2' zstandard==0.24.0 polars==1.32.0 duckdb==1.3.2 pylance==0.39.0 cramjam boto3 pytest~=7.0 py4j==0.10.9.9 requests parameterized==0.9.0 'daft>=0.7.6' 'datafusion>=54,<55' datasketches if python -c "import sys; sys.exit(0 if sys.version_info >= (3, 12) else 1)"; then python -m pip install pyarrow==24.0.0 numpy==2.4.6 pandas==2.3.3 flake8==7.1.2 else diff --git a/paimon-python/dev/lint-python.sh b/paimon-python/dev/lint-python.sh index 79dc14ede959..d81860735d79 100755 --- a/paimon-python/dev/lint-python.sh +++ b/paimon-python/dev/lint-python.sh @@ -198,7 +198,7 @@ function pytest_check() { # 3.6/3.7 run a curated core subset (their dep ceiling rules out the # vector/index/multimodal/blob suites); 3.10+ run the full suite. if [ "$PYTHON_VERSION" = "3.6" ] || [ "$PYTHON_VERSION" = "3.7" ]; then - TEST_DIR="pypaimon/tests/py36 pypaimon/tests/file_io_test.py" + TEST_DIR="pypaimon/tests/py36 pypaimon/tests/file_io_test.py pypaimon/tests/s3_atomic_write_test.py" echo "Running core test subset for Python $PYTHON_VERSION: $TEST_DIR" else TEST_DIR="pypaimon/tests pypaimon/acceptance --ignore=pypaimon/tests/py36 --ignore=pypaimon/tests/e2e --ignore=pypaimon/tests/torch_read_test.py" diff --git a/paimon-python/dev/requirements.txt b/paimon-python/dev/requirements.txt index dc9b4b4e911a..5f5132d1010a 100644 --- a/paimon-python/dev/requirements.txt +++ b/paimon-python/dev/requirements.txt @@ -40,5 +40,6 @@ requests>=2.21.0,<3 urllib3>=1.26,<3 zstandard>=0.19,<1 backports.zstd>=1.0.0,<1.4.0; python_version >= "3.9" and python_version < "3.14" +boto3>=1.23.10,<2 cramjam>=1.3.0,<3; python_version>="3.7" pyyaml>=5.4,<7 diff --git a/paimon-python/pypaimon/filesystem/pyarrow_file_io.py b/paimon-python/pypaimon/filesystem/pyarrow_file_io.py index 362994fb65ef..571260cb8ad1 100644 --- a/paimon-python/pypaimon/filesystem/pyarrow_file_io.py +++ b/paimon-python/pypaimon/filesystem/pyarrow_file_io.py @@ -62,12 +62,14 @@ def __init__(self, path: str, catalog_options: Options): scheme, netloc, _ = self.parse_location(path) self.uri_reader_factory = UriReaderFactory(catalog_options) self._is_oss = scheme in {"oss"} + self._is_s3 = scheme in {"s3", "s3a", "s3n"} self._oss_bucket = None _oss_impl = self.properties.get(OssOptions.OSS_IMPL) self._use_jindo = False self._legacy_bucket_checked = False self._legacy_bucket_error = None self._legacy_bucket_lock = threading.Lock() + self._s3_atomic_client = None if self._is_oss: self._oss_bucket = self._extract_oss_bucket(path) @@ -98,11 +100,13 @@ def __getstate__(self): state = self.__dict__.copy() # threading.Lock cannot be pickled; recreated in __setstate__. state.pop("_legacy_bucket_lock", None) + state.pop("_s3_atomic_client", None) return state def __setstate__(self, state): self.__dict__.update(state) self._legacy_bucket_lock = threading.Lock() + self._s3_atomic_client = None @staticmethod def parse_location(location: str): @@ -578,6 +582,29 @@ def delete_directory_quietly(self, directory: str): self.logger.warning(f"Exception occurs when deleting directory {directory}", exc_info=True) def try_to_write_atomic(self, path: str, content: str) -> bool: + if self._is_s3: + from botocore.exceptions import ClientError + + uri = urlparse(path) + if uri.scheme and uri.scheme not in {"s3", "s3a", "s3n"}: + raise ValueError(f"Object path uses a different filesystem: {path}") + path_str = self.to_filesystem_path(path) + bucket, _, key = path_str.partition("/") + if not bucket or not key or key == ".": + raise ValueError(f"Invalid S3 object path: {path}") + if self._get_file_info(path_str).type == pafs.FileType.Directory: + return False + + client = self._get_s3_atomic_client() + try: + client.put_object(Bucket=bucket, Key=key, Body=content.encode("utf-8")) + return True + except ClientError as error: + if error.response.get("Error", {}).get("Code") in ( + "PreconditionFailed", "ConditionalRequestConflict"): + return False + raise + if self.exists(path): path_str = self.to_filesystem_path(path) file_info = self._get_file_info(path_str) @@ -594,6 +621,47 @@ def try_to_write_atomic(self, path: str, content: str) -> bool: self.delete_quietly(temp_path) return success + def _get_s3_atomic_client(self): + if self._s3_atomic_client is not None: + return self._s3_atomic_client + + import boto3 + from botocore.config import Config + + access_key = self._get_property( + S3Options.S3_ACCESS_KEY_ID.key(), + *self._s3_key_variants("access-key", "access.key")) + secret_key = self._get_property( + S3Options.S3_ACCESS_KEY_SECRET.key(), + *self._s3_key_variants("secret-key", "secret.key")) + session_token = self._get_property( + S3Options.S3_SECURITY_TOKEN.key(), + *self._s3_key_variants("session-token", "session.token", + "security-token", "security.token")) + endpoint = self._get_s3_property("endpoint", S3Options.S3_ENDPOINT.key()) + region = self._get_s3_property("region", S3Options.S3_REGION.key()) + path_style = (self._get_s3_boolean_property("path-style-access") or + self._get_s3_boolean_property("path.style.access")) + addressing_style = "path" if path_style else ( + "virtual" if self._pyarrow_gte_16 else "auto") + + if endpoint and "://" not in endpoint: + endpoint = "https://" + endpoint + config_kwargs = {"s3": {"addressing_style": addressing_style}} + if "request_checksum_calculation" in Config.OPTION_DEFAULTS: + config_kwargs["request_checksum_calculation"] = "when_required" + client = boto3.session.Session().client( + "s3", endpoint_url=endpoint, aws_access_key_id=access_key, + aws_secret_access_key=secret_key, aws_session_token=session_token, + region_name=region, config=Config(**config_kwargs)) + client.meta.events.register( + "before-sign.s3.PutObject", self._add_atomic_write_header) + self._s3_atomic_client = client + return client + + def _add_atomic_write_header(self, request, **kwargs): + request.headers["If-None-Match"] = "*" + def copy_file(self, source_path: str, target_path: str, overwrite: bool = False): if not overwrite and self.exists(target_path): raise FileExistsError(f"Target file {target_path} already exists and overwrite=False") diff --git a/paimon-python/pypaimon/tests/s3_atomic_write_test.py b/paimon-python/pypaimon/tests/s3_atomic_write_test.py new file mode 100644 index 000000000000..2186217bfef6 --- /dev/null +++ b/paimon-python/pypaimon/tests/s3_atomic_write_test.py @@ -0,0 +1,163 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. + +import os +import socketserver +import threading +import unittest +from http.server import BaseHTTPRequestHandler, HTTPServer +from unittest import mock +from urllib.parse import urlsplit + +import pyarrow.fs as pafs + +try: + import boto3 + from botocore.exceptions import ClientError +except ImportError: + boto3 = None + ClientError = Exception + +from pypaimon.common.options import Options +from pypaimon.common.options.config import S3Options +from pypaimon.filesystem.pyarrow_file_io import PyArrowFileIO + + +class _Server(socketserver.ThreadingMixIn, HTTPServer): + daemon_threads = True + + +class _ConditionalPutHandler(BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + + def _respond(self, status, body=b""): + self.send_response(status) + self.send_header("Content-Type", "application/xml") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + + def do_PUT(self): + body = self.rfile.read(int(self.headers["Content-Length"])) + if self.server.deny_put: + return self._respond(403, b"AccessDenied") + _, bucket, key = urlsplit(self.path).path.split("/", 2) + with self.server.lock: + old = self.server.objects.get((bucket, key)) + if self.headers.get("If-None-Match") == "*" and old is not None: + return self._respond(412, b"PreconditionFailed") + self.server.objects[(bucket, key)] = body + self.server.headers_seen.append(dict(self.headers)) + self._respond(200) + + def log_message(self, *args): + pass + + +@unittest.skipUnless(boto3 is not None, "requires boto3") +class S3AtomicWriteTest(unittest.TestCase): + def setUp(self): + self.server = _Server(("127.0.0.1", 0), _ConditionalPutHandler) + self.server.objects = {} + self.server.headers_seen = [] + self.server.deny_put = False + self.server.lock = threading.Lock() + self.thread = threading.Thread(target=self.server.serve_forever) + self.thread.start() + self.endpoint = "http://127.0.0.1:{}".format(self.server.server_port) + self.env = mock.patch.dict(os.environ, { + "NO_PROXY": "127.0.0.1,localhost", + "no_proxy": "127.0.0.1,localhost", + }) + self.env.start() + + def tearDown(self): + self.env.stop() + self.server.shutdown() + self.server.server_close() + self.thread.join() + + def _s3_io(self): + options = Options({ + S3Options.S3_ACCESS_KEY_ID.key(): "ak", + S3Options.S3_ACCESS_KEY_SECRET.key(): "sk", + S3Options.S3_ENDPOINT.key(): self.endpoint, + S3Options.S3_REGION.key(): "us-east-1", + "fs.s3.path.style.access": "true", + }) + with mock.patch.object(PyArrowFileIO, "_initialize_s3_fs", return_value=mock.Mock()): + io = PyArrowFileIO("s3://test-bucket/", options) + io.filesystem = mock.Mock(spec=pafs.S3FileSystem) + io.filesystem.get_file_info.return_value = [ + pafs.FileInfo("test-bucket/table", pafs.FileType.NotFound)] + return io + + def test_concurrent_s3_writers_publish_only_one_value(self): + io = self._s3_io() + results = [] + barrier = threading.Barrier(2) + + def write(value): + barrier.wait() + results.append(io.try_to_write_atomic( + "s3://test-bucket/table/snapshot/snapshot-1", value)) + + threads = [threading.Thread(target=write, args=(value,)) + for value in ("first", "second")] + for thread in threads: + thread.start() + for thread in threads: + thread.join(10) + + self.assertFalse(any(thread.is_alive() for thread in threads)) + self.assertCountEqual([True, False], results) + self.assertIn(self.server.objects[("test-bucket", "table/snapshot/snapshot-1")], + (b"first", b"second")) + self.assertEqual("*", self.server.headers_seen[0]["If-None-Match"]) + self.assertIn("if-none-match", self.server.headers_seen[0]["Authorization"]) + + def test_cross_bucket_s3_write_uses_target_bucket(self): + io = self._s3_io() + io.try_to_write_atomic("s3://other-bucket/table/schema/schema-0", "data") + + self.assertEqual({("other-bucket", "table/schema/schema-0"): b"data"}, + self.server.objects) + + def test_wrong_uri_scheme_does_not_write(self): + io = self._s3_io() + with self.assertRaisesRegex(ValueError, "different filesystem"): + io.try_to_write_atomic("oss://test-bucket/table/schema/schema-0", "bad") + self.assertEqual({}, self.server.objects) + + def test_write_permission_error_does_not_look_like_conflict(self): + io = self._s3_io() + self.server.deny_put = True + + with self.assertRaises(ClientError): + io.try_to_write_atomic( + "s3://test-bucket/table/snapshot/snapshot-1", "data") + + self.assertEqual({}, self.server.objects) + + def test_atomic_write_rejects_directory(self): + io = self._s3_io() + io.filesystem.get_file_info.return_value = [ + pafs.FileInfo("test-bucket/table", pafs.FileType.Directory)] + + self.assertFalse(io.try_to_write_atomic("s3://test-bucket/table", "data")) + + self.assertEqual({}, self.server.objects)