Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
78 changes: 48 additions & 30 deletions python/lightning_sdk/api/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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:
Comment thread
codexceed marked this conversation as resolved.
"""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.
Expand Down
94 changes: 88 additions & 6 deletions python/tests/api/test_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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"},
Expand Down Expand Up @@ -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()
Expand Down
99 changes: 99 additions & 0 deletions python/tests/cli/utils/test_logging.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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
Loading