diff --git a/s3proxy/handlers/base.py b/s3proxy/handlers/base.py
index 43277f4..be69923 100644
--- a/s3proxy/handlers/base.py
+++ b/s3proxy/handlers/base.py
@@ -29,6 +29,7 @@
from ..config import Settings
from ..errors import S3Error, raise_for_client_error
from ..state import MultipartStateManager
+from ..state.complete_lock import CompleteUploadLock, create_complete_upload_lock
from ..utils import etag_matches, parse_http_date
logger: BoundLogger = structlog.get_logger(__name__)
@@ -140,10 +141,12 @@ def __init__(
settings: Settings,
credentials_store: dict[str, str],
multipart_manager: MultipartStateManager,
+ complete_upload_lock: CompleteUploadLock | None = None,
):
self.settings = settings
self.credentials_store = credentials_store
self.multipart_manager = multipart_manager
+ self.complete_upload_lock = complete_upload_lock or create_complete_upload_lock()
self.keyring = settings.keyring
def _client(self, creds: S3Credentials) -> S3Client:
diff --git a/s3proxy/handlers/multipart/lifecycle.py b/s3proxy/handlers/multipart/lifecycle.py
index ecfd264..7e0901e 100644
--- a/s3proxy/handlers/multipart/lifecycle.py
+++ b/s3proxy/handlers/multipart/lifecycle.py
@@ -23,6 +23,7 @@
MultipartUploadState,
PartMetadata,
delete_upload_state,
+ load_multipart_metadata,
persist_upload_state,
plaintext_attr_cache,
save_multipart_metadata,
@@ -151,107 +152,158 @@ async def handle_complete_multipart_upload(
bucket, key = self._parse_path(request.url.path)
async with self._client(creds) as client:
upload_id, _ = self._extract_multipart_params(request)
-
- state = await self.multipart_manager.complete_upload(bucket, key, upload_id)
- if not state:
- state = await self._recover_upload_state(
- client, bucket, key, upload_id, context="for complete"
+ async with self.complete_upload_lock.hold(bucket, key, upload_id):
+ return await self._handle_complete_multipart_upload_locked(
+ request, creds, client, bucket, key, upload_id
)
- if state.deferred_copy_tail:
- logger.info(
- "COMPLETE_MULTIPART_DEFERRED_TAIL_PENDING",
- bucket=bucket,
- key=key,
- upload_id=upload_id[:20] + "...",
- tail_bytes=len(state.deferred_copy_tail),
- )
- state = await self._flush_deferred_copy_tail_for_complete(
- client, bucket, key, upload_id, state
- )
-
- # Parse client's part list
- body = await request.body()
- client_parts = self._parse_client_parts(body)
+ async def _handle_complete_multipart_upload_locked(
+ self,
+ request: Request,
+ creds: S3Credentials,
+ client: S3Client,
+ bucket: str,
+ key: str,
+ upload_id: str,
+ ) -> Response:
+ idempotent = await self._try_idempotent_complete_response(client, bucket, key, upload_id)
+ if idempotent is not None:
+ return idempotent
- # Build S3 parts list
- s3_parts, completed_parts, total_plaintext = self._build_s3_parts(
- client_parts, state, bucket, key, upload_id
+ state = await self.multipart_manager.complete_upload(bucket, key, upload_id)
+ if not state:
+ state = await self._recover_upload_state(
+ client, bucket, key, upload_id, context="for complete"
)
+ if state.deferred_copy_tail:
logger.info(
- "COMPLETE_MULTIPART",
+ "COMPLETE_MULTIPART_DEFERRED_TAIL_PENDING",
bucket=bucket,
key=key,
upload_id=upload_id[:20] + "...",
- client_parts=len(completed_parts),
- s3_parts=len(s3_parts),
- total_mb=f"{total_plaintext / 1024 / 1024:.2f}MB",
+ tail_bytes=len(state.deferred_copy_tail),
+ )
+ state = await self._flush_deferred_copy_tail_for_complete(
+ client, bucket, key, upload_id, state
)
- # Complete in S3
- try:
- complete_resp = await self._complete_multipart_upload_with_retry(
- client, bucket, key, upload_id, s3_parts, completed_parts
- )
- except ClientError as e:
- await self._handle_complete_error(
- e, client, bucket, key, upload_id, s3_parts, completed_parts, total_plaintext
- )
- else:
- plaintext_attr_cache.put(
- bucket,
- key,
- str(complete_resp.get("ETag", "")).strip('"'),
- total_plaintext,
- synthetic_multipart_etag(total_plaintext),
- )
+ # Parse client's part list
+ body = await request.body()
+ client_parts = self._parse_client_parts(body)
- # Save metadata first, then delete state.
- # Order matters: if metadata save fails, state is preserved
- # so the upload can be retried. Deleting state first would
- # lose the DEK, making the object permanently undecryptable.
- # Prefer the kid recorded when the upload was created; if the state
- # predates it (e.g. older recovered state), fall back to the
- # completing credential's key.
- if state.kid:
- kid, kek = state.kid, self.keyring.key_by_id(state.kid)
- else:
- kid, kek = self.keyring.key_for(creds.access_key)
- wrapped_dek = crypto.wrap_key(state.dek, kek)
- await save_multipart_metadata(
- client,
+ # Build S3 parts list
+ s3_parts, completed_parts, total_plaintext = self._build_s3_parts(
+ client_parts, state, bucket, key, upload_id
+ )
+
+ logger.info(
+ "COMPLETE_MULTIPART",
+ bucket=bucket,
+ key=key,
+ upload_id=upload_id[:20] + "...",
+ client_parts=len(completed_parts),
+ s3_parts=len(s3_parts),
+ total_mb=f"{total_plaintext / 1024 / 1024:.2f}MB",
+ )
+
+ # Complete in S3
+ try:
+ complete_resp = await self._complete_multipart_upload_with_retry(
+ client, bucket, key, upload_id, s3_parts, completed_parts
+ )
+ except ClientError as e:
+ await self._handle_complete_error(
+ e, client, bucket, key, upload_id, s3_parts, completed_parts, total_plaintext
+ )
+ else:
+ plaintext_attr_cache.put(
bucket,
key,
- MultipartMetadata(
- version=2,
- part_count=len(completed_parts),
- total_plaintext_size=total_plaintext,
- parts=completed_parts,
- wrapped_dek=wrapped_dek,
- kid=kid,
- ),
+ str(complete_resp.get("ETag", "")).strip('"'),
+ total_plaintext,
+ synthetic_multipart_etag(total_plaintext),
)
- await delete_upload_state(client, bucket, key, upload_id)
- logger.info(
- "COMPLETE_MULTIPART_SUCCESS",
- bucket=bucket,
- key=key,
- upload_id=upload_id[:20] + "...",
- total_parts=len(completed_parts),
- total_mb=f"{total_plaintext / 1024 / 1024:.2f}MB",
- )
+ # Save metadata first, then delete state.
+ # Order matters: if metadata save fails, state is preserved
+ # so the upload can be retried. Deleting state first would
+ # lose the DEK, making the object permanently undecryptable.
+ # Prefer the kid recorded when the upload was created; if the state
+ # predates it (e.g. older recovered state), fall back to the
+ # completing credential's key.
+ if state.kid:
+ kid, kek = state.kid, self.keyring.key_by_id(state.kid)
+ else:
+ kid, kek = self.keyring.key_for(creds.access_key)
+ wrapped_dek = crypto.wrap_key(state.dek, kek)
+ await save_multipart_metadata(
+ client,
+ bucket,
+ key,
+ MultipartMetadata(
+ version=2,
+ part_count=len(completed_parts),
+ total_plaintext_size=total_plaintext,
+ parts=completed_parts,
+ wrapped_dek=wrapped_dek,
+ kid=kid,
+ ),
+ )
+ await delete_upload_state(client, bucket, key, upload_id)
- location = f"{self.settings.s3_endpoint}/{bucket}/{key}"
- etag = hashlib.md5(
- str(state.total_plaintext_size).encode(), usedforsecurity=False
- ).hexdigest()
+ logger.info(
+ "COMPLETE_MULTIPART_SUCCESS",
+ bucket=bucket,
+ key=key,
+ upload_id=upload_id[:20] + "...",
+ total_parts=len(completed_parts),
+ total_mb=f"{total_plaintext / 1024 / 1024:.2f}MB",
+ )
- return Response(
- content=xml_responses.complete_multipart(location, bucket, key, etag),
- media_type="application/xml",
- )
+ location = f"{self.settings.s3_endpoint}/{bucket}/{key}"
+ etag = hashlib.md5(
+ str(state.total_plaintext_size).encode(), usedforsecurity=False
+ ).hexdigest()
+
+ return Response(
+ content=xml_responses.complete_multipart(location, bucket, key, etag),
+ media_type="application/xml",
+ )
+
+ async def _try_idempotent_complete_response(
+ self, client: S3Client, bucket: str, key: str, upload_id: str
+ ) -> Response | None:
+ """Return success if a peer pod already finished this upload."""
+ meta = await load_multipart_metadata(client, bucket, key)
+ if meta is None:
+ return None
+
+ try:
+ head = await client.head_object(bucket, key)
+ except ClientError:
+ return None
+
+ expected_ciphertext_size = sum(p.ciphertext_size for p in meta.parts)
+ if head.get("ContentLength") != expected_ciphertext_size:
+ return None
+
+ logger.info(
+ "COMPLETE_MULTIPART_IDEMPOTENT",
+ bucket=bucket,
+ key=key,
+ upload_id=upload_id[:20] + "...",
+ total_mb=f"{meta.total_plaintext_size / 1024 / 1024:.2f}MB",
+ )
+
+ location = f"{self.settings.s3_endpoint}/{bucket}/{key}"
+ etag = hashlib.md5(
+ str(meta.total_plaintext_size).encode(), usedforsecurity=False
+ ).hexdigest()
+ return Response(
+ content=xml_responses.complete_multipart(location, bucket, key, etag),
+ media_type="application/xml",
+ )
def _parse_client_parts(self, body: bytes) -> list[dict]:
client_parts = []
diff --git a/s3proxy/state/__init__.py b/s3proxy/state/__init__.py
index 2ad3520..3bf991f 100644
--- a/s3proxy/state/__init__.py
+++ b/s3proxy/state/__init__.py
@@ -11,6 +11,9 @@
# Plaintext attribute cache for listings
from .attr_cache import PlaintextAttrCache, plaintext_attr_cache, synthetic_multipart_etag
+# Complete upload serialization (HA)
+from .complete_lock import CompleteUploadLock, create_complete_upload_lock
+
# State manager and storage
from .manager import MAX_INTERNAL_PARTS_PER_CLIENT, MultipartStateManager
@@ -90,6 +93,9 @@
"synthetic_multipart_etag",
# Recovery
"reconstruct_upload_state_from_s3",
+ # Complete lock
+ "CompleteUploadLock",
+ "create_complete_upload_lock",
# Serialization
"json_dumps",
"json_loads",
diff --git a/s3proxy/state/complete_lock.py b/s3proxy/state/complete_lock.py
new file mode 100644
index 0000000..dec0538
--- /dev/null
+++ b/s3proxy/state/complete_lock.py
@@ -0,0 +1,180 @@
+"""Distributed lock for CompleteMultipartUpload.
+
+HA deployments load-balance across many pods. Without a per-upload lock, two pods
+can both call upstream CompleteMultipartUpload for the same upload_id (one after
+recovering state the other already deleted), which surfaces as "multipart
+completion is already in progress" and can leave the client with NoSuchUpload.
+"""
+
+from __future__ import annotations
+
+import asyncio
+import os
+import uuid
+from collections.abc import AsyncIterator
+from contextlib import asynccontextmanager
+from typing import TYPE_CHECKING
+
+import structlog
+from structlog.stdlib import BoundLogger
+
+from ..errors import S3Error
+
+if TYPE_CHECKING:
+ from redis.asyncio import Redis
+
+logger: BoundLogger = structlog.get_logger(__name__)
+
+COMPLETE_LOCK_TTL_SECONDS = int(os.environ.get("S3PROXY_COMPLETE_LOCK_TTL_SECONDS", "7200"))
+COMPLETE_LOCK_ACQUIRE_TIMEOUT_SECONDS = float(
+ os.environ.get("S3PROXY_COMPLETE_LOCK_ACQUIRE_TIMEOUT_SECONDS", "7200")
+)
+COMPLETE_LOCK_POLL_INTERVAL_SECONDS = float(
+ os.environ.get("S3PROXY_COMPLETE_LOCK_POLL_INTERVAL_SECONDS", "0.25")
+)
+
+_REDIS_PREFIX = "s3proxy:complete-lock:"
+
+
+class CompleteUploadLock:
+ """Serialize CompleteMultipartUpload per (bucket, key, upload_id)."""
+
+ def __init__(
+ self,
+ redis_client: Redis | None = None,
+ *,
+ ttl_seconds: int = COMPLETE_LOCK_TTL_SECONDS,
+ acquire_timeout_seconds: float = COMPLETE_LOCK_ACQUIRE_TIMEOUT_SECONDS,
+ poll_interval_seconds: float = COMPLETE_LOCK_POLL_INTERVAL_SECONDS,
+ ) -> None:
+ self._redis = redis_client
+ self._ttl = ttl_seconds
+ self._acquire_timeout = acquire_timeout_seconds
+ self._poll_interval = poll_interval_seconds
+ self._memory_locks: dict[str, asyncio.Lock] = {}
+ self._memory_guard = asyncio.Lock()
+
+ def _storage_key(self, bucket: str, key: str, upload_id: str) -> str:
+ return f"{bucket}:{key}:{upload_id}"
+
+ def _redis_key(self, bucket: str, key: str, upload_id: str) -> str:
+ return f"{_REDIS_PREFIX}{bucket}:{key}:{upload_id}"
+
+ @asynccontextmanager
+ async def hold(self, bucket: str, key: str, upload_id: str) -> AsyncIterator[None]:
+ if self._redis is not None:
+ async with self._redis_hold(bucket, key, upload_id):
+ yield
+ else:
+ async with self._memory_hold(bucket, key, upload_id):
+ yield
+
+ @asynccontextmanager
+ async def _memory_hold(self, bucket: str, key: str, upload_id: str) -> AsyncIterator[None]:
+ lk = self._storage_key(bucket, key, upload_id)
+ async with self._memory_guard:
+ lock = self._memory_locks.get(lk)
+ if lock is None:
+ lock = asyncio.Lock()
+ self._memory_locks[lk] = lock
+
+ await lock.acquire()
+ logger.debug(
+ "COMPLETE_LOCK_ACQUIRED",
+ bucket=bucket,
+ key=key,
+ upload_id=upload_id[:20] + "..." if len(upload_id) > 20 else upload_id,
+ backend="memory",
+ )
+ try:
+ yield
+ finally:
+ lock.release()
+ logger.debug(
+ "COMPLETE_LOCK_RELEASED",
+ bucket=bucket,
+ key=key,
+ upload_id=upload_id[:20] + "..." if len(upload_id) > 20 else upload_id,
+ backend="memory",
+ )
+
+ @asynccontextmanager
+ async def _redis_hold(self, bucket: str, key: str, upload_id: str) -> AsyncIterator[None]:
+ redis_key = self._redis_key(bucket, key, upload_id)
+ token = uuid.uuid4().hex
+ deadline = asyncio.get_running_loop().time() + self._acquire_timeout
+
+ while True:
+ acquired = await self._redis.set(redis_key, token, nx=True, ex=self._ttl)
+ if acquired:
+ logger.debug(
+ "COMPLETE_LOCK_ACQUIRED",
+ bucket=bucket,
+ key=key,
+ upload_id=upload_id[:20] + "..." if len(upload_id) > 20 else upload_id,
+ backend="redis",
+ )
+ break
+
+ if asyncio.get_running_loop().time() >= deadline:
+ logger.warning(
+ "COMPLETE_LOCK_ACQUIRE_TIMEOUT",
+ bucket=bucket,
+ key=key,
+ upload_id=upload_id[:20] + "..." if len(upload_id) > 20 else upload_id,
+ timeout_seconds=self._acquire_timeout,
+ )
+ raise S3Error.slow_down(
+ "CompleteMultipartUpload is in progress on another instance; retry later"
+ )
+
+ await asyncio.sleep(self._poll_interval)
+
+ try:
+ yield
+ finally:
+ await self._release_redis_lock(redis_key, token)
+ logger.debug(
+ "COMPLETE_LOCK_RELEASED",
+ bucket=bucket,
+ key=key,
+ upload_id=upload_id[:20] + "..." if len(upload_id) > 20 else upload_id,
+ backend="redis",
+ )
+
+ async def _release_redis_lock(self, redis_key: str, token: str) -> None:
+ import redis.asyncio as redis
+
+ from .storage import MAX_WATCH_RETRIES, WATCH_RETRY_BASE_DELAY_SEC
+
+ for attempt in range(MAX_WATCH_RETRIES):
+ async with self._redis.pipeline(transaction=True) as pipe:
+ try:
+ await pipe.watch(redis_key)
+ current = await self._redis.get(redis_key)
+ if current is None:
+ await pipe.unwatch()
+ return
+ current_token = current.decode() if isinstance(current, bytes) else current
+ if current_token != token:
+ await pipe.unwatch()
+ return
+ pipe.multi()
+ pipe.delete(redis_key)
+ await pipe.execute()
+ return
+ except redis.WatchError:
+ if attempt == MAX_WATCH_RETRIES - 1:
+ logger.warning(
+ "COMPLETE_LOCK_RELEASE_CONFLICT",
+ redis_key=redis_key,
+ )
+ return
+ await asyncio.sleep(WATCH_RETRY_BASE_DELAY_SEC * (2**attempt))
+
+
+def create_complete_upload_lock() -> CompleteUploadLock:
+ """Build a lock using Redis when the HA client is initialized."""
+ from .redis import _redis_client
+
+ return CompleteUploadLock(redis_client=_redis_client)
diff --git a/tests/unit/test_complete_upload_lock.py b/tests/unit/test_complete_upload_lock.py
new file mode 100644
index 0000000..9e6c758
--- /dev/null
+++ b/tests/unit/test_complete_upload_lock.py
@@ -0,0 +1,285 @@
+"""CompleteMultipartUpload cross-pod serialization.
+
+Prod incident 2026-07-31: HAProxy routed concurrent CompleteMultipartUpload
+requests for the same upload_id to different pods. Each pod deleted/recovered
+state independently and both called upstream CompleteMultipartUpload, yielding
+"This multipart completion is already in progress" and eventually NoSuchUpload.
+"""
+
+from __future__ import annotations
+
+import asyncio
+from urllib.parse import urlencode
+
+import pytest
+
+from s3proxy import crypto
+from s3proxy.state import (
+ CompleteUploadLock,
+ MultipartMetadata,
+ PartMetadata,
+ save_multipart_metadata,
+)
+from s3proxy.state.metadata import persist_upload_state
+
+_INTERNAL_FOR_CLIENT = {1: 1, 2: 21}
+
+
+class _FakeURL:
+ def __init__(self, path: str, query: str) -> None:
+ self.path = path
+ self.query = query
+
+
+class _FakeRequest:
+ def __init__(self, path: str, query: str, body: bytes) -> None:
+ self.url = _FakeURL(path, query)
+ self.headers: dict[str, str] = {}
+ self._body = body
+
+ async def body(self) -> bytes:
+ return self._body
+
+
+class _ClientCM:
+ def __init__(self, client) -> None:
+ self._client = client
+
+ async def __aenter__(self):
+ return self._client
+
+ async def __aexit__(self, *exc) -> bool:
+ return False
+
+
+def _complete_request(
+ bucket: str, key: str, upload_id: str, client_parts: list[int]
+) -> _FakeRequest:
+ parts_xml = "".join(
+ f"{n}"etag-{n}""
+ for n in client_parts
+ )
+ body = f"{parts_xml}".encode()
+ return _FakeRequest(f"/{bucket}/{key}", urlencode({"uploadId": upload_id}), body)
+
+
+async def _seed_upload(mock_s3, dek: bytes, bucket: str, key: str, chunks: dict[int, bytes]) -> str:
+ resp = await mock_s3.create_multipart_upload(bucket, key)
+ upload_id = resp["UploadId"]
+ for client_part, plaintext in chunks.items():
+ internal = _INTERNAL_FOR_CLIENT[client_part]
+ nonce = crypto.derive_part_nonce(upload_id, internal)
+ ciphertext = crypto.encrypt(plaintext, dek, nonce)
+ await mock_s3.upload_part(bucket, key, upload_id, internal, ciphertext)
+ return upload_id
+
+
+class _CountingCompleteClient:
+ """Tracks upstream complete calls and optionally blocks the first one."""
+
+ def __init__(self, inner, *, block_first_until: asyncio.Event | None = None) -> None:
+ self._inner = inner
+ self.complete_calls = 0
+ self._block_first_until = block_first_until
+ self._first_complete_started = asyncio.Event()
+
+ async def complete_multipart_upload(self, bucket, key, upload_id, parts):
+ self.complete_calls += 1
+ if self.complete_calls == 1:
+ self._first_complete_started.set()
+ if self._block_first_until is not None:
+ await self._block_first_until.wait()
+ return await self._inner.complete_multipart_upload(bucket, key, upload_id, parts)
+
+ def __getattr__(self, name):
+ return getattr(self._inner, name)
+
+
+@pytest.fixture
+def complete_upload_lock(mock_redis):
+ return CompleteUploadLock(
+ redis_client=mock_redis,
+ ttl_seconds=30,
+ acquire_timeout_seconds=5,
+ poll_interval_seconds=0.05,
+ )
+
+
+@pytest.fixture
+def handler_with_lock(settings, manager, complete_upload_lock):
+ from s3proxy.handlers.multipart import MultipartHandlerMixin
+
+ h = MultipartHandlerMixin(settings, {}, manager, complete_upload_lock=complete_upload_lock)
+ return h
+
+
+@pytest.mark.asyncio
+async def test_complete_lock_serializes_concurrent_upstream_calls(
+ handler_with_lock, mock_s3, mock_s3_client, settings, credentials
+):
+ bucket, key = "test-bucket", "backup/large.db"
+ kid, kek = settings.keyring.key_for(credentials.access_key)
+ dek = crypto.generate_dek()
+ chunk1 = b"first-part" * 64
+ chunk2 = b"second-part" * 32
+ upload_id = await _seed_upload(mock_s3, dek, bucket, key, {1: chunk1, 2: chunk2})
+ await persist_upload_state(
+ mock_s3_client, bucket, key, upload_id, crypto.wrap_key(dek, kek), kid
+ )
+
+ await handler_with_lock.multipart_manager.create_upload(bucket, key, upload_id, dek, kid)
+ for client_part, plaintext in ((1, chunk1), (2, chunk2)):
+ internal = _INTERNAL_FOR_CLIENT[client_part]
+ nonce = crypto.derive_part_nonce(upload_id, internal)
+ ciphertext = crypto.encrypt(plaintext, dek, nonce)
+ from s3proxy.state import InternalPartMetadata, PartMetadata
+
+ await handler_with_lock.multipart_manager.add_part(
+ bucket,
+ key,
+ upload_id,
+ PartMetadata(
+ part_number=client_part,
+ plaintext_size=len(plaintext),
+ ciphertext_size=len(ciphertext),
+ etag="etag",
+ md5="etag",
+ internal_parts=[
+ InternalPartMetadata(
+ internal_part_number=internal,
+ plaintext_size=len(plaintext),
+ ciphertext_size=len(ciphertext),
+ etag="etag",
+ )
+ ],
+ ),
+ )
+
+ release_first = asyncio.Event()
+ counting_client = _CountingCompleteClient(mock_s3_client, block_first_until=release_first)
+ handler_with_lock._client = lambda creds: _ClientCM(counting_client)
+
+ req = _complete_request(bucket, key, upload_id, [1, 2])
+
+ first = asyncio.create_task(
+ handler_with_lock.handle_complete_multipart_upload(req, credentials)
+ )
+ await counting_client._first_complete_started.wait()
+
+ second = asyncio.create_task(
+ handler_with_lock.handle_complete_multipart_upload(req, credentials)
+ )
+ await asyncio.sleep(0.1)
+ assert counting_client.complete_calls == 1
+
+ release_first.set()
+ await asyncio.gather(first, second)
+
+ assert counting_client.complete_calls == 1
+ assert f"{bucket}/{key}" in mock_s3.objects
+
+
+@pytest.mark.asyncio
+async def test_complete_lock_idempotent_when_peer_already_finished(
+ handler_with_lock, mock_s3, mock_s3_client, settings, credentials
+):
+ bucket, key = "test-bucket", "backup/done.db"
+ kid, kek = settings.keyring.key_for(credentials.access_key)
+ dek = crypto.generate_dek()
+ chunk1 = b"payload-a" * 32
+ chunk2 = b"payload-b" * 16
+ upload_id = await _seed_upload(mock_s3, dek, bucket, key, {1: chunk1, 2: chunk2})
+
+ ct1 = crypto.encrypt(chunk1, dek, crypto.derive_part_nonce(upload_id, 1))
+ ct2 = crypto.encrypt(chunk2, dek, crypto.derive_part_nonce(upload_id, 21))
+ mock_s3.objects[f"{bucket}/{key}"] = {
+ "Body": ct1 + ct2,
+ "Metadata": {},
+ "ContentType": "application/octet-stream",
+ "ContentLength": len(ct1) + len(ct2),
+ "ETag": "done",
+ "LastModified": __import__("datetime").datetime.now(__import__("datetime").UTC),
+ }
+
+ parts = [
+ PartMetadata(
+ part_number=1,
+ plaintext_size=len(chunk1),
+ ciphertext_size=len(ct1),
+ etag="e1",
+ md5="e1",
+ ),
+ PartMetadata(
+ part_number=2,
+ plaintext_size=len(chunk2),
+ ciphertext_size=len(ct2),
+ etag="e2",
+ md5="e2",
+ ),
+ ]
+ await save_multipart_metadata(
+ mock_s3_client,
+ bucket,
+ key,
+ MultipartMetadata(
+ version=2,
+ part_count=2,
+ total_plaintext_size=len(chunk1) + len(chunk2),
+ parts=parts,
+ wrapped_dek=crypto.wrap_key(dek, kek),
+ kid=kid,
+ ),
+ )
+
+ counting_client = _CountingCompleteClient(mock_s3_client)
+ handler_with_lock._client = lambda creds: _ClientCM(counting_client)
+
+ resp = await handler_with_lock.handle_complete_multipart_upload(
+ _complete_request(bucket, key, upload_id, [1, 2]), credentials
+ )
+
+ assert resp.status_code == 200
+ assert counting_client.complete_calls == 0
+
+
+@pytest.mark.asyncio
+async def test_complete_lock_memory_backend_serializes():
+ lock = CompleteUploadLock(
+ redis_client=None,
+ acquire_timeout_seconds=2,
+ poll_interval_seconds=0.01,
+ )
+ order: list[str] = []
+ gate = asyncio.Event()
+
+ async def worker(name: str) -> None:
+ async with lock.hold("b", "k", "upload-1"):
+ order.append(f"{name}-start")
+ if name == "a":
+ gate.set()
+ await asyncio.sleep(0.05)
+ order.append(f"{name}-end")
+
+ task_a = asyncio.create_task(worker("a"))
+ await gate.wait()
+ task_b = asyncio.create_task(worker("b"))
+ await asyncio.gather(task_a, task_b)
+
+ assert order == ["a-start", "a-end", "b-start", "b-end"]
+
+
+@pytest.mark.asyncio
+async def test_complete_lock_redis_acquire_timeout_raises(mock_redis):
+ lock = CompleteUploadLock(
+ redis_client=mock_redis,
+ ttl_seconds=30,
+ acquire_timeout_seconds=0.15,
+ poll_interval_seconds=0.05,
+ )
+
+ async with lock.hold("bucket", "key", "upload-locked"):
+ with pytest.raises(Exception) as exc:
+ async with lock.hold("bucket", "key", "upload-locked"):
+ pass
+
+ assert getattr(exc.value, "code", None) == "SlowDown"