diff --git a/README.md b/README.md index ef05c1e8..49a545c7 100644 --- a/README.md +++ b/README.md @@ -50,6 +50,7 @@ uploaded = bucket.upload_resumable( "videos/automatic.mp4", video, content_type="video/mp4", + on_progress=lambda uploaded, total: print(f"{uploaded}/{total}"), ) print(uploaded.name) upload_session = bucket.create_upload_session( @@ -110,7 +111,9 @@ part count, and expiration time for a resumable upload. `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. +part or progress callback fails, it makes a best-effort abort and raises the +original error. `on_progress` runs after each successful part with cumulative +uploaded bytes and the total size. `upload_part()` returns immutable part metadata and can safely retry the same part number to replace that part. `locks.get()` returns immutable lock availability, expiry, and fencing-token diff --git a/src/volcano_sdk/storage.py b/src/volcano_sdk/storage.py index 11330a6c..1237cf7b 100644 --- a/src/volcano_sdk/storage.py +++ b/src/volcano_sdk/storage.py @@ -5,7 +5,7 @@ import base64 import binascii import json -from collections.abc import Generator, Mapping, Sequence +from collections.abc import Callable, Generator, Mapping, Sequence from contextlib import contextmanager, suppress from dataclasses import dataclass from datetime import datetime @@ -564,6 +564,7 @@ def upload_resumable( *, content_type: str = "application/octet-stream", part_size: int | None = None, + on_progress: Callable[[int, int], None] | None = None, ) -> StorageObject: """Upload bytes or a binary stream through a resumable session.""" path = _storage_path(path) @@ -577,7 +578,13 @@ def upload_resumable( ) upload_succeeded = False try: - self._upload_session_parts(path, source, session) + self._upload_session_parts( + path, + source, + session, + total_size, + on_progress, + ) upload_succeeded = True finally: if not upload_succeeded: @@ -589,14 +596,21 @@ def _upload_session_parts( path: str, source: BinaryIO, session: UploadSession, + total_size: int, + on_progress: Callable[[int, int], None] | None, ) -> None: + uploaded = 0 for part_index in range(session.total_parts): + part = _read_upload_part(source, session.part_size) self.upload_part( path, session_id=session.session_id, part_number=part_index + 1, - data=_read_upload_part(source, session.part_size), + data=part, ) + uploaded += len(part) + if on_progress is not None: + on_progress(uploaded, total_size) def _abort_failed_upload(self, path: str, session_id: str) -> None: with suppress(Exception): diff --git a/tests/unit/test_facade.py b/tests/unit/test_facade.py index c1185b76..4db45c52 100644 --- a/tests/unit/test_facade.py +++ b/tests/unit/test_facade.py @@ -816,6 +816,47 @@ def test_storage_uploads_bytes_with_server_selected_chunks() -> None: assert object_.name == "videos/demo.mp4" +def test_storage_reports_progress_after_each_uploaded_part() -> 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") + progress: list[tuple[int, int]] = [] + + client.storage.from_("assets").upload_resumable( + "file.bin", + b"abcdefghij", + on_progress=lambda uploaded, total: progress.append((uploaded, total)), + ) + + assert progress == [(4, 10), (8, 10), (10, 10)] + + +def test_storage_aborts_when_a_progress_callback_fails() -> 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") + + def fail_progress(_uploaded: int, _total: int) -> None: + message = "progress failed" + raise RuntimeError(message) + + with pytest.raises(RuntimeError, match="progress failed"): + client.storage.from_("assets").upload_resumable( + "file.bin", + b"abcdefgh", + on_progress=fail_progress, + ) + + assert [operation for operation, _ in transport.calls[-2:]] == [ + "uploadPart", + "abortUploadSession", + ] + + def test_storage_aborts_after_part_failure_without_masking_the_error() -> None: transport = FakeTransport() transport.upload_session_part_size = 4