diff --git a/README.md b/README.md index ee9726b5..393b303f 100644 --- a/README.md +++ b/README.md @@ -103,9 +103,10 @@ object; `public_url` is set only when the object is public. Pass an HTTP byte range to download only part of an object. `create_upload_session()` returns the immutable server-selected part size, part count, and expiration time for a resumable upload. -`upload_resumable()` creates a session, chunks the bytes using the -server-selected part size, and completes the upload. If a part fails, it makes -a best-effort abort and raises the original error. +`upload_resumable()` accepts bytes or a binary file-like object, creates a +session, and uploads server-sized chunks. It streams seekable files directly; +non-seekable inputs are spooled to a temporary file with bounded reads. If a +part fails, it makes a best-effort abort and raises the original error. `upload_part()` returns immutable part metadata and can safely retry the same part number to replace that part. `get_upload_session()` returns immutable progress and uploaded-part metadata for diff --git a/src/volcano_sdk/storage.py b/src/volcano_sdk/storage.py index 1292421c..11330a6c 100644 --- a/src/volcano_sdk/storage.py +++ b/src/volcano_sdk/storage.py @@ -5,11 +5,13 @@ import base64 import binascii import json -from collections.abc import Mapping, Sequence -from contextlib import suppress +from collections.abc import Generator, Mapping, Sequence +from contextlib import contextmanager, suppress from dataclasses import dataclass from datetime import datetime -from typing import Any, Protocol, cast +from io import SEEK_END, BytesIO +from tempfile import TemporaryFile +from typing import Any, BinaryIO, Protocol, cast from urllib.parse import quote from ._transport import ( @@ -21,7 +23,6 @@ invoke, response_payload, ) -from .errors import VolcanoError from .models import ( JSONValue, StorageObject, @@ -40,6 +41,8 @@ _INVALID_PUBLIC_URL_PATH = "Public URL paths cannot contain dot segments" _JWT_PART_COUNT = 3 _HTTP_PARTIAL_CONTENT = 206 +_UPLOAD_SPOOL_READ_SIZE = 1_048_576 +_UPLOAD_SOURCE_UNAVAILABLE = "Upload source is temporarily unavailable" def _optional_datetime(value: object) -> datetime | None: @@ -198,6 +201,65 @@ def _encoded_storage_path(path: str) -> str: return "/".join(quote(segment, safe="") for segment in segments) +def _remaining_upload_bytes(source: BinaryIO) -> int | None: + try: + if not source.seekable(): + return None + position = source.tell() + except (AttributeError, OSError, ValueError): + return None + try: + try: + source.seek(0, SEEK_END) + remaining = max(0, source.tell() - position) + except (OSError, ValueError): + remaining = None + finally: + source.seek(position) + return remaining + + +def _spool_upload_source(source: BinaryIO, target: BinaryIO) -> None: + while True: + chunk = cast("bytes | None", source.read(_UPLOAD_SPOOL_READ_SIZE)) + if chunk is None: + raise BlockingIOError(_UPLOAD_SOURCE_UNAVAILABLE) + if chunk == b"": + return + target.write(chunk) + + +def _read_upload_part(source: BinaryIO, part_size: int) -> bytes: + part = bytearray() + while len(part) < part_size: + chunk = cast("bytes | None", source.read(part_size - len(part))) + if chunk is None: + raise BlockingIOError(_UPLOAD_SOURCE_UNAVAILABLE) + if chunk == b"": + break + part.extend(chunk) + return bytes(part) + + +@contextmanager +def _resumable_upload_source( + data: bytes | BinaryIO, +) -> Generator[tuple[BinaryIO, int], None, None]: + if isinstance(data, bytes): + with BytesIO(data) as source: + yield source, len(data) + return + remaining = _remaining_upload_bytes(data) + if remaining is not None: + yield data, remaining + return + with TemporaryFile(mode="w+b") as source: + _spool_upload_source(data, source) + total_size = source.tell() + source.seek(0) + yield source, total_size + + class StorageContext(Protocol): """Client capabilities required by object storage.""" @@ -498,42 +560,46 @@ def abort_upload_session( def upload_resumable( self, path: str, - data: bytes, + data: bytes | BinaryIO, *, content_type: str = "application/octet-stream", part_size: int | None = None, ) -> StorageObject: - """Upload bytes through a server-managed resumable session.""" - session = self.create_upload_session( - path, - total_size=len(data), - content_type=content_type, - part_size=part_size, - ) - try: - self._upload_session_parts(path, data, session) - except VolcanoError: - self._abort_failed_upload(path, session.session_id) - raise - return self.complete_upload_session(path, session_id=session.session_id) + """Upload bytes or a binary stream through a resumable session.""" + path = _storage_path(path) + self._client._session_token() + with _resumable_upload_source(data) as (source, total_size): + session = self.create_upload_session( + path, + total_size=total_size, + content_type=content_type, + part_size=part_size, + ) + upload_succeeded = False + try: + self._upload_session_parts(path, source, session) + upload_succeeded = True + finally: + if not upload_succeeded: + self._abort_failed_upload(path, session.session_id) + return self.complete_upload_session(path, session_id=session.session_id) def _upload_session_parts( self, path: str, - data: bytes, + source: BinaryIO, session: UploadSession, ) -> None: for part_index in range(session.total_parts): - offset = part_index * session.part_size self.upload_part( path, session_id=session.session_id, part_number=part_index + 1, - data=data[offset : offset + session.part_size], + data=_read_upload_part(source, session.part_size), ) def _abort_failed_upload(self, path: str, session_id: str) -> None: - with suppress(VolcanoError): + with suppress(Exception): self.abort_upload_session(path, session_id=session_id) def list( diff --git a/tests/unit/test_facade.py b/tests/unit/test_facade.py index 5870e8c0..79e830b6 100644 --- a/tests/unit/test_facade.py +++ b/tests/unit/test_facade.py @@ -4,7 +4,8 @@ import json from dataclasses import dataclass from datetime import UTC, datetime -from typing import Any +from io import SEEK_END, BytesIO +from typing import Any, BinaryIO, cast import pytest @@ -39,6 +40,7 @@ def __init__(self) -> None: self.upload_session_total_parts = 3 self.fail_upload_part_number: int | None = None self.fail_abort_upload = False + self.raise_abort_error = False def auth_signin(self, **kwargs: Any) -> FakeResponse: self.calls.append(("authSignin", kwargs)) @@ -148,6 +150,9 @@ def get_upload_session(self, **kwargs: Any) -> FakeResponse: def abort_upload_session(self, **kwargs: Any) -> FakeResponse: self.calls.append(("abortUploadSession", kwargs)) + if self.raise_abort_error: + msg = "abort transport failed" + raise RuntimeError(msg) if self.fail_abort_upload: return FakeResponse(500, {"error": "abort failed"}) return FakeResponse(200, {"message": "upload session aborted"}) @@ -239,6 +244,91 @@ def release_project_lock(self, **kwargs: Any) -> FakeResponse: return FakeResponse(204) +class BoundedBytesIO(BytesIO): + def __init__(self, value: bytes) -> None: + super().__init__(value) + self.read_sizes: list[int] = [] + + def read(self, size: int | None = -1) -> bytes: + if size is None or size < 0: + msg = "unbounded read" + raise RuntimeError(msg) + self.read_sizes.append(size) + return super().read(size) + + +class ShortReadBytesIO(BoundedBytesIO): + def read(self, size: int | None = -1) -> bytes: + if size is None or size < 0: + msg = "unbounded read" + raise RuntimeError(msg) + return super().read(min(size, 2)) + + +class BoundedNonSeekableReader: + def __init__(self, value: bytes) -> None: + self._value = value + self._offset = 0 + self.read_sizes: list[int] = [] + + def read(self, size: int = -1) -> bytes: + if size < 0: + msg = "unbounded read" + raise RuntimeError(msg) + self.read_sizes.append(size) + chunk = self._value[self._offset : self._offset + size] + self._offset += len(chunk) + return chunk + + def seekable(self) -> bool: + return False + + +class TemporarilyUnavailableReader: + def __init__(self) -> None: + self.read_sizes: list[int] = [] + + def read(self, size: int = -1) -> bytes | None: + if not self.read_sizes: + self.read_sizes.append(size) + return None + return b"" + + def seekable(self) -> bool: + return False + + +class ReadOnlyStream: + def __init__(self, value: bytes) -> None: + self._source = BytesIO(value) + + def read(self, size: int = -1) -> bytes: + return self._source.read(size) + + +class FailingSeekableReader(BoundedBytesIO): + def read(self, size: int | None = -1) -> bytes: + if self.tell() >= 4: + msg = "reader failed" + raise RuntimeError(msg) + return super().read(size) + + +class RestoreFailingBytesIO(BytesIO): + def __init__(self, value: bytes) -> None: + super().__init__(value) + self._end_was_probed = False + + def seek(self, offset: int, whence: int = 0) -> int: + if self._end_was_probed and whence == 0: + msg = "restore failed" + raise OSError(msg) + position = super().seek(offset, whence) + if whence == SEEK_END: + self._end_was_probed = True + return position + + def anon_key_with_project_id(project_id: str | None) -> str: payload = {} if project_id is None else {"project_id": project_id} encoded = base64.urlsafe_b64encode(json.dumps(payload).encode()).rstrip(b"=") @@ -647,6 +737,193 @@ def test_storage_aborts_after_part_failure_without_masking_the_error() -> None: ] +def test_storage_streams_seekable_uploads_with_server_selected_reads() -> None: + transport = FakeTransport() + transport.upload_session_part_size = 4 + transport.upload_session_total_parts = 3 + client = VolcanoClient(anon_key="anon-key", _transport=transport) + client.auth.sign_in(email="user@example.com", password="secret") + source = BoundedBytesIO(b"abcdefghij") + + client.storage.from_("assets").upload_resumable("file.bin", source) + + assert source.read_sizes + assert max(source.read_sizes) <= 4 + upload_calls = [call for call in transport.calls if call[0] == "uploadPart"] + assert [call[1]["request"].data for call in upload_calls] == [ + b"abcd", + b"efgh", + b"ij", + ] + + +def test_storage_clamps_a_seekable_source_positioned_past_eof() -> None: + transport = FakeTransport() + transport.upload_session_total_parts = 0 + client = VolcanoClient(anon_key="anon-key", _transport=transport) + client.auth.sign_in(email="user@example.com", password="secret") + source = BytesIO(b"a") + source.seek(2) + + client.storage.from_("assets").upload_resumable("file.bin", source) + + create_call = next( + call for call in transport.calls if call[0] == "createUploadSession" + ) + assert create_call[1]["request"].total_size == 0 + + +def test_storage_surfaces_a_failed_seekable_position_restore() -> None: + transport = FakeTransport() + client = VolcanoClient(anon_key="anon-key", _transport=transport) + client.auth.sign_in(email="user@example.com", password="secret") + + with pytest.raises(OSError, match="restore failed"): + client.storage.from_("assets").upload_resumable( + "file.bin", + RestoreFailingBytesIO(b"abcdefgh"), + ) + + assert all(operation != "createUploadSession" for operation, _ in transport.calls) + + +def test_storage_fills_parts_when_a_seekable_source_returns_short_reads() -> None: + transport = FakeTransport() + transport.upload_session_part_size = 4 + transport.upload_session_total_parts = 3 + client = VolcanoClient(anon_key="anon-key", _transport=transport) + client.auth.sign_in(email="user@example.com", password="secret") + + client.storage.from_("assets").upload_resumable( + "file.bin", + ShortReadBytesIO(b"abcdefghij"), + ) + + upload_calls = [call for call in transport.calls if call[0] == "uploadPart"] + assert [call[1]["request"].data for call in upload_calls] == [ + b"abcd", + b"efgh", + b"ij", + ] + + +def test_storage_spools_non_seekable_uploads_with_bounded_reads() -> None: + transport = FakeTransport() + transport.upload_session_part_size = 4 + transport.upload_session_total_parts = 3 + client = VolcanoClient(anon_key="anon-key", _transport=transport) + client.auth.sign_in(email="user@example.com", password="secret") + source = BoundedNonSeekableReader(b"abcdefghij") + + client.storage.from_("assets").upload_resumable( + "file.bin", + cast("BinaryIO", source), + ) + + assert source.read_sizes + assert max(source.read_sizes) <= 1_048_576 + upload_calls = [call for call in transport.calls if call[0] == "uploadPart"] + assert [call[1]["request"].data for call in upload_calls] == [ + b"abcd", + b"efgh", + b"ij", + ] + + +def test_storage_spools_read_only_streams_without_a_seekability_probe() -> None: + transport = FakeTransport() + transport.upload_session_part_size = 4 + transport.upload_session_total_parts = 2 + client = VolcanoClient(anon_key="anon-key", _transport=transport) + client.auth.sign_in(email="user@example.com", password="secret") + + client.storage.from_("assets").upload_resumable( + "file.bin", + cast("BinaryIO", ReadOnlyStream(b"abcdefgh")), + ) + + upload_calls = [call for call in transport.calls if call[0] == "uploadPart"] + assert [call[1]["request"].data for call in upload_calls] == [b"abcd", b"efgh"] + + +def test_storage_aborts_when_a_stream_reader_raises_an_unexpected_error() -> None: + transport = FakeTransport() + transport.upload_session_part_size = 4 + transport.upload_session_total_parts = 2 + client = VolcanoClient(anon_key="anon-key", _transport=transport) + client.auth.sign_in(email="user@example.com", password="secret") + + with pytest.raises(RuntimeError, match="reader failed"): + client.storage.from_("assets").upload_resumable( + "file.bin", + FailingSeekableReader(b"abcdefgh"), + ) + + assert [operation for operation, _ in transport.calls[-2:]] == [ + "uploadPart", + "abortUploadSession", + ] + + +def test_storage_preserves_reader_error_when_abort_cleanup_raises() -> None: + transport = FakeTransport() + transport.upload_session_part_size = 4 + transport.upload_session_total_parts = 2 + transport.raise_abort_error = True + client = VolcanoClient(anon_key="anon-key", _transport=transport) + client.auth.sign_in(email="user@example.com", password="secret") + + with pytest.raises(RuntimeError, match="reader failed"): + client.storage.from_("assets").upload_resumable( + "file.bin", + FailingSeekableReader(b"abcdefgh"), + ) + + +def test_storage_rejects_temporarily_unavailable_nonblocking_sources() -> None: + transport = FakeTransport() + client = VolcanoClient(anon_key="anon-key", _transport=transport) + client.auth.sign_in(email="user@example.com", password="secret") + source = TemporarilyUnavailableReader() + + with pytest.raises(BlockingIOError, match="temporarily unavailable"): + client.storage.from_("assets").upload_resumable( + "file.bin", + cast("BinaryIO", source), + ) + + assert all(operation != "createUploadSession" for operation, _ in transport.calls) + + +def test_storage_validates_authentication_before_spooling() -> None: + transport = FakeTransport() + client = VolcanoClient(anon_key="anon-key", _transport=transport) + source = BoundedNonSeekableReader(b"abcdefghij") + + with pytest.raises(RuntimeError, match="active session"): + client.storage.from_("assets").upload_resumable( + "file.bin", + cast("BinaryIO", source), + ) + + assert source.read_sizes == [] + + +def test_storage_validates_path_before_spooling() -> None: + transport = FakeTransport() + client = VolcanoClient(anon_key="anon-key", _transport=transport) + client.auth.sign_in(email="user@example.com", password="secret") + source = BoundedNonSeekableReader(b"abcdefghij") + + with pytest.raises(ValueError, match="non-empty string"): + client.storage.from_("assets").upload_resumable( + "", + cast("BinaryIO", source), + ) + + assert source.read_sizes == [] + + @pytest.mark.parametrize("invalid_paths", [[], [""], b"abc"]) def test_storage_remove_rejects_invalid_paths_before_transport( invalid_paths: Any,