diff --git a/python/lightning_sdk/api/utils.py b/python/lightning_sdk/api/utils.py index 2b312595..471eb884 100644 --- a/python/lightning_sdk/api/utils.py +++ b/python/lightning_sdk/api/utils.py @@ -119,6 +119,8 @@ def __init__(self) -> None: _MAX_WORKERS = 10 # S3, R2, and GCS all cap multipart uploads at 10,000 parts _MAX_UPLOAD_PARTS = 10000 +# 5xx statuses that reject the request itself rather than signalling a transient fault +_NON_RETRYABLE_UPLOAD_STATUSES = (501, 505) def _local_file_matches_size(local_path: str, expected_size: Optional[int]) -> bool: @@ -378,39 +380,55 @@ def _single_part_upload(self) -> None: if self.extra_headers: headers.update(self.extra_headers) - with open(self.local_path, "rb") as f: + with contextlib.ExitStack() as stack: + handle = stack.enter_context(open(self.local_path, "rb")) + body: Any = handle if self.show_progress: - with tqdm.wrapattr( - f, - "read", - desc=f"Uploading {os.path.split(self.local_path)[1]}", - total=self.filesize, - unit="B", - unit_scale=True, - unit_divisor=1024, - ) as wrapped_file: - r = requests.put( - urls[0]["url"], - data=_IterableFileWrapper(wrapped_file), - timeout=30, - headers=headers or None, + body = _IterableFileWrapper( + stack.enter_context( + tqdm.wrapattr( + handle, + "read", + desc=f"Uploading {os.path.split(self.local_path)[1]}", + total=self.filesize, + unit="B", + unit_scale=True, + unit_divisor=1024, + ) ) - else: - r = requests.put(urls[0]["url"], data=f, timeout=30, headers=headers or None) - - if r.status_code != 200: - # Retry transient server errors / throttling, and also 401/403 from - # storage: every attempt signs a fresh URL, which heals expired - # signatures and newly issued storage credentials that haven't - # propagated yet (e.g. right after a lightning_storage folder is - # created). The backoff decorator only retries HTTPError/ - # RequestException, so raise an HTTPError for these instead of - # failing immediately. - if r.status_code >= 500 or r.status_code in (401, 403, 429): - raise HTTPError( - f"Transient error uploading file '{self.local_path}'. Status code: {r.status_code}", response=r ) - raise RuntimeError(f"Failed to upload file '{self.local_path}'. Status code: {r.status_code}") + if self.filesize == 0: + # requests infers length by truth test, so an empty stream would go out + # "Transfer-Encoding: chunked", which presigned PUTs reject with 501. + body = b"" + r = requests.put(urls[0]["url"], data=body, timeout=30, headers=headers or None) + + self._raise_for_upload_status(r) + + def _raise_for_upload_status(self, r: requests.Response) -> None: + """Turn a non-200 storage PUT response into the right exception type. + + Args: + r: The response from the presigned storage PUT. + + Raises: + HTTPError: On statuses worth retrying, so the backoff-wrapped caller + signs a fresh URL and tries again. + RuntimeError: On any other non-200 status. + """ + if r.status_code == 200: + return + # 401/403 are worth retrying here because each attempt signs a fresh URL, which + # heals expired signatures and storage credentials that haven't propagated yet. + # Backoff only retries HTTPError/RequestException, hence the type split. + retryable = r.status_code not in _NON_RETRYABLE_UPLOAD_STATUSES and ( + r.status_code >= 500 or r.status_code in (401, 403, 429) + ) + if retryable: + raise HTTPError( + f"Transient error uploading file '{self.local_path}'. Status code: {r.status_code}", response=r + ) + raise RuntimeError(f"Failed to upload file '{self.local_path}'. Status code: {r.status_code}") def _upload_part_with_recovery(self, url_info: Dict[str, Any], upload_id: str) -> Dict[str, Any]: """Upload one part, falling back to re-signing its URL on failure. diff --git a/python/tests/api/test_utils.py b/python/tests/api/test_utils.py index 6e46624a..cc085711 100644 --- a/python/tests/api/test_utils.py +++ b/python/tests/api/test_utils.py @@ -30,6 +30,8 @@ from lightning_sdk.machine import Machine _TEST_ENDPOINT_BASE = "https://api.example.com/v1/projects/test-project-id/artifacts" +# Tests that prepare a real request need a URL requests can parse, not a bare placeholder. +_TEST_SIGNED_PUT_URL = "https://storage.example.com/signed-put-url" def _make_mocked_blob_uploader(monkeypatch, file_path, remote_path, **kwargs): @@ -291,6 +293,83 @@ def test_blob_uploader_single_part_no_completion(tmp_path, monkeypatch): put_mock.assert_called_once() +def _capturing_put(captured): + """A ``requests.put`` stand-in that records the headers requests would really send. + + The other upload tests assert on the arguments handed to a ``Mock``, which never builds + a request and so cannot see headers requests derives from the body. Preparing it here is + what makes the length assertions meaningful, and why these tests need a parseable URL. + """ + + def put(url, data=None, headers=None, **_): + captured.append(requests.Request("PUT", url, data=data, headers=headers).prepare()) + return Mock(status_code=200) + + return put + + +@pytest.mark.parametrize("progress_bar", [False, True]) +@pytest.mark.parametrize(("content", "expected_length"), [(b"", "0"), (b"print('hi')\n", "12")]) +def test_blob_uploader_single_part_declares_body_length(tmp_path, monkeypatch, progress_bar, content, expected_length): + """Single-part uploads declare their length instead of streaming chunked. + + requests falls back to chunked transfer encoding when it cannot determine a body's + length, and a zero-length body is indistinguishable from an unknown one by its truth + value. Presigned storage PUTs answer chunked uploads with 501. + """ + file_path = tmp_path / "__init__.py" + file_path.write_bytes(content) + + uploader = _make_mocked_blob_uploader(monkeypatch, file_path=str(file_path), remote_path="app/__init__.py") + uploader.show_progress = progress_bar + + captured = [] + monkeypatch.setattr( + lightning_sdk.api.utils.requests, + "post", + Mock(return_value=_blob_upload_response("app/__init__.py", "", [{"url": _TEST_SIGNED_PUT_URL}])), + ) + monkeypatch.setattr(lightning_sdk.api.utils.requests, "put", _capturing_put(captured)) + + uploader() + + assert len(captured) == 1 + assert captured[0].headers["Content-Length"] == expected_length + assert "Transfer-Encoding" not in captured[0].headers + + +def test_blob_uploader_single_part_empty_file_keeps_signed_headers(tmp_path, monkeypatch): + """The empty body still carries the presigned and extra headers the upload needs.""" + file_path = tmp_path / "empty.bin" + file_path.write_bytes(b"") + + uploader = _make_mocked_blob_uploader( + monkeypatch, + file_path=str(file_path), + remote_path="remote-path", + content_type="text/plain", + extra_headers={"x-ms-blob-type": "BlockBlob"}, + ) + + captured = [] + monkeypatch.setattr( + lightning_sdk.api.utils.requests, + "post", + Mock( + return_value=_blob_upload_response( + "remote-path", "", [{"url": _TEST_SIGNED_PUT_URL, "headers": {"Content-Type": "text/plain"}}] + ) + ), + ) + monkeypatch.setattr(lightning_sdk.api.utils.requests, "put", _capturing_put(captured)) + + uploader() + + assert captured[0].headers["Content-Type"] == "text/plain" + assert captured[0].headers["x-ms-blob-type"] == "BlockBlob" + assert captured[0].headers["Content-Length"] == "0" + + def _make_mocked_model_uploader(monkeypatch, file_path, remote_path): # Threadpools don't like mocks as input, so we just use a regular map here monkeypatch.setattr(lightning_sdk.api.utils.ThreadPoolExecutor, "map", map) @@ -711,9 +790,9 @@ def test_resolve_path_mappings(): assert path_mappings[1].connection_path == "" -def _make_single_part_blob_uploader(tmp_path, progress_bar=False, extra_headers=None): +def _make_single_part_blob_uploader(tmp_path, progress_bar=False, extra_headers=None, content=b"hello world"): file_path = tmp_path / "test_file.txt" - file_path.write_text("hello world") + file_path.write_bytes(content) with mock.patch( "lightning_sdk.api.utils._authenticate_and_get_auth_headers", return_value={"Authorization": "Bearer test-token"}, @@ -826,14 +905,17 @@ def test_single_part_uploader_retries_auth_errors_with_fresh_url(mock_requests, assert mock_requests.post.call_count == 2 +# 404 rejects the request, 501 rejects the protocol: neither improves on a retry, and the +# empty-body path has to reach the same verdict as the streaming one. +@pytest.mark.parametrize("status_code", [404, 501]) +@pytest.mark.parametrize("content", [b"hello world", b""]) @mock.patch("time.sleep", return_value=None) @mock.patch("lightning_sdk.api.utils.requests") -def test_single_part_uploader_does_not_retry_client_error(mock_requests, _, tmp_path): - uploader = _make_single_part_blob_uploader(tmp_path) +def test_single_part_uploader_does_not_retry_client_error(mock_requests, _, tmp_path, content, status_code): + uploader = _make_single_part_blob_uploader(tmp_path, content=content) _configure_single_part_create(mock_requests) - # A non-transient 4xx should fail immediately without retrying. - mock_requests.put.return_value = Mock(status_code=404) + mock_requests.put.return_value = Mock(status_code=status_code) with pytest.raises(RuntimeError, match="Failed to upload file"): uploader() diff --git a/python/tests/cli/utils/test_logging.py b/python/tests/cli/utils/test_logging.py index 72eebb82..446de4e4 100644 --- a/python/tests/cli/utils/test_logging.py +++ b/python/tests/cli/utils/test_logging.py @@ -9,12 +9,17 @@ """ import contextlib +import os +import pathlib +import subprocess import sys from unittest import mock import click import click.testing +import pytest +import lightning_sdk from lightning_sdk.__version__ import __version__ from lightning_sdk.cli.utils.logging import ( CommandLoggingGroup, @@ -471,3 +476,97 @@ def subcommand(): assert result.exit_code == 0 assert "Hello" in result.output + + +# The tests above prove _notify_exception renders a clean panel and that logging_excepthook calls +# it. Neither says the hook is ever reached: Click formats ClickException itself and lets every +# other error escape, so the panel depends on an ordinary error travelling all the way out of +# main() to the interpreter. These cover that last link, which is invisible to CliRunner because +# the runner catches exceptions before the interpreter ever sees them. + +# Rejected before any network call, so it exercises the failure path offline. +_FAILING_COMMAND = ["cp", "lit:///a/b", "lit:///c/d"] + + +class TestErrorsReachTheExcepthook: + """A plain error from a command has to reach sys.excepthook for the CLI panel to appear.""" + + def test_running_a_command_installs_the_excepthook(self, monkeypatch): + """The group callback points sys.excepthook at the handler that prints the panel.""" + from lightning_sdk.cli.entrypoint import main_cli + + monkeypatch.setattr(sys, "excepthook", sys.__excepthook__) + monkeypatch.setattr("lightning_sdk.cli.utils.logging._log_command", mock.Mock()) + + click.testing.CliRunner().invoke(main_cli, _FAILING_COMMAND, prog_name="lightning") + + assert sys.excepthook is logging_excepthook + + def test_click_lets_a_plain_error_escape(self, monkeypatch): + """Click formats only ClickException, so an SDK error must come out of main() unchanged.""" + from lightning_sdk.cli.entrypoint import main_cli + + monkeypatch.setattr(sys, "excepthook", sys.__excepthook__) + monkeypatch.setattr("lightning_sdk.cli.utils.logging._log_command", mock.Mock()) + + with pytest.raises(ValueError, match="Cannot copy between two remote URLs"): + main_cli.main(_FAILING_COMMAND, prog_name="lightning", standalone_mode=True) + + +# Driving a real interpreter is the only way to see what a user sees: the hook runs after the +# exception has left every frame a test could wrap. _log_command is started rather than scoped +# for the same reason - a with block would unwind before the hook fires, letting it hit the network. +# The driver lives in a temp dir, so sys.path[0] is that dir and an installed copy of the package +# would win over this checkout; _PACKAGE_ROOT pins the subprocess to the same one pytest imported. +_CLI_FAILURE_DRIVER = """ +import sys +from unittest import mock + +mock.patch("lightning_sdk.cli.utils.logging._log_command").start() +from lightning_sdk.cli.entrypoint import main_cli + +sys.argv = ["lightning", *{args!r}] +main_cli() +""" + + +_PACKAGE_ROOT = str(pathlib.Path(lightning_sdk.__file__).resolve().parents[1]) + + +def _run_cli_driver(tmp_path, debug): + driver = tmp_path / "driver.py" + driver.write_text(_CLI_FAILURE_DRIVER.format(args=_FAILING_COMMAND)) + result = subprocess.run( + [sys.executable, str(driver)], + capture_output=True, + text=True, + env={ + **os.environ, + "PYTHONPATH": _PACKAGE_ROOT, + "LIGHTNING_API_KEY": "dummy-api-key", + "COLUMNS": "200", + "LIGHTNING_DEBUG": debug, + }, + check=False, + ) + return result.returncode, result.stdout + result.stderr + + +def test_cli_reports_errors_as_a_panel_not_a_traceback(tmp_path): + """End to end: a failing command prints the error panel and no traceback.""" + returncode, output = _run_cli_driver(tmp_path, debug="") + + assert returncode == 1 + assert "Lightning CLI Error" in output + assert "Cannot copy between two remote URLs" in output + assert "Traceback (most recent call last)" not in output + assert "LIGHTNING_DEBUG=1" in output + + +def test_cli_shows_the_traceback_when_debug_is_set(tmp_path): + """LIGHTNING_DEBUG=1 is the documented way back to the traceback, so it has to work.""" + returncode, output = _run_cli_driver(tmp_path, debug="1") + + assert returncode == 1 + assert "Full traceback" in output + assert "Cannot copy between two remote URLs" in output