diff --git a/go/job.go b/go/job.go index fd07374e8..3ea547f7f 100644 --- a/go/job.go +++ b/go/job.go @@ -40,6 +40,8 @@ type Job struct { artifactsSource string artifactsDestination string + maxRunAttempts int64 + currentRunAttempt int64 } // JobDict is the JSON-friendly public representation of a job. @@ -103,6 +105,7 @@ type jobOptions struct { artifactsSource string artifactsDest string maxRuntime int + maxRunAttempts int64 scratchDisks []ScratchDisk } @@ -154,6 +157,9 @@ type JobOptions struct { ArtifactsDestination string // MaxRuntime limits the job runtime in seconds. MaxRuntime int + // MaxRunAttempts is the max number of run attempts for this job. + // 0 means unset (backend default). 1 means a single attempt (no retries). + MaxRunAttempts int64 // ScratchDisks mounts temporary scratch storage for studio-backed jobs. ScratchDisks []ScratchDisk } @@ -306,6 +312,24 @@ func (j *Job) ArtifactsDestination() string { return j.artifactsDestination } +// MaxRunAttempts returns the max number of run attempts for this job. +// 0 means unset. +func (j *Job) MaxRunAttempts() int64 { + if j == nil { + return 0 + } + return j.maxRunAttempts +} + +// CurrentRunAttempt returns the current run attempt for this job. +// 0 means unset. +func (j *Job) CurrentRunAttempt() int64 { + if j == nil { + return 0 + } + return j.currentRunAttempt +} + // GetJob returns an existing job by name or ID. func GetJob(name string, opts ...JobOptions) (*Job, error) { resolved := applyJobOptions(opts...) @@ -640,6 +664,7 @@ func applyJobOptions(opts ...JobOptions) jobOptions { resolved.artifactsSource = opts[0].ArtifactsSource resolved.artifactsDest = opts[0].ArtifactsDestination resolved.maxRuntime = opts[0].MaxRuntime + resolved.maxRunAttempts = opts[0].MaxRunAttempts resolved.scratchDisks = opts[0].ScratchDisks } return resolved @@ -735,6 +760,8 @@ func jobFromModel(model *models.V1Job, opts jobOptions) *Job { result.studioID = model.Spec.CloudspaceID result.artifactsSource = model.Spec.ArtifactsSource result.artifactsDestination = model.Spec.ArtifactsDestination + result.maxRunAttempts = model.Spec.MaxRunAttempts + result.currentRunAttempt = model.Spec.CurrentRunAttempt } return result } @@ -767,6 +794,7 @@ func jobSpec(machine, command string, opts jobOptions) *models.V1JobSpec { ImageSecretRef: opts.imageCredentials, InstanceName: machine, PathMappings: jobPathMappings(opts), + MaxRunAttempts: opts.maxRunAttempts, RequestedRunDurationSeconds: maxRuntime(opts.maxRuntime), Spot: opts.interruptible, Volumes: scratchVolumes(opts.scratchDisks), diff --git a/go/job_test.go b/go/job_test.go index c163c6be8..527f51a1f 100644 --- a/go/job_test.go +++ b/go/job_test.go @@ -38,6 +38,33 @@ func TestJobExposesStartAndStopTimes(t *testing.T) { "StoppedAt = %v, want zero time", existing.StoppedAt()) } +func TestJobExposesRunAttemptFields(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]any{ + "id": "job-retry", + "name": "train", + "projectId": "project-1", + "state": "pending", + "spec": map[string]any{ + "maxRunAttempts": 3, + "currentRunAttempt": 2, + "parentJobId": "job-parent", + }, + }) + })) + defer server.Close() + t.Setenv("LIGHTNING_CLOUD_URL", server.URL) + + existing, err := lit.GetJob("train", lit.JobOptions{Teamspace: mustTeamspace(t, "project-1", "default", "alice")}) + require.NoErrorf(t, err, + "GetJob returned error") + assert.Falsef(t, existing.MaxRunAttempts() != 3, + "MaxRunAttempts = %d, want 3", existing.MaxRunAttempts()) + assert.Falsef(t, existing.CurrentRunAttempt() != 2, + "CurrentRunAttempt = %d, want 2", existing.CurrentRunAttempt()) +} + func TestJobGetWithIDUsesSimpleStruct(t *testing.T) { j, err := lit.GetJob("train", lit.JobOptions{ID: "job-1", Teamspace: mustTeamspace(t, "project-1", "")}) require.NoErrorf(t, err, @@ -288,6 +315,7 @@ func TestJobRunMapsAdvancedV2Options(t *testing.T) { ImageClusterCredentials bool `json:"imageClusterCredentials"` ImageSecretRef string `json:"imageSecretRef"` RequestedRunDurationSeconds string `json:"requestedRunDurationSeconds"` + MaxRunAttempts int64 `json:"maxRunAttempts"` PathMappings []struct { ContainerPath string `json:"containerPath"` ConnectionName string `json:"connectionName"` @@ -308,6 +336,8 @@ func TestJobRunMapsAdvancedV2Options(t *testing.T) { "imageSecretRef = %q, want docker-secret", body.Spec.ImageSecretRef) assert.Falsef(t, body.Spec.RequestedRunDurationSeconds != "7200", "requestedRunDurationSeconds = %q, want 7200", body.Spec.RequestedRunDurationSeconds) + assert.Falsef(t, body.Spec.MaxRunAttempts != 3, + "maxRunAttempts = %d, want 3", body.Spec.MaxRunAttempts) assert.Falsef(t, len(body.Spec.PathMappings) != 2, "pathMappings length = %d, want 2", len(body.Spec.PathMappings)) @@ -330,6 +360,9 @@ func TestJobRunMapsAdvancedV2Options(t *testing.T) { "image": "registry.example/train:latest", "artifactsSource": "/outputs", "artifactsDestination": "efs:data:outputs/run-1", + "maxRunAttempts": 3, + "currentRunAttempt": 1, + "parentJobId": "", }, }) })) @@ -349,6 +382,7 @@ func TestJobRunMapsAdvancedV2Options(t *testing.T) { ArtifactsSource: "/outputs", ArtifactsDestination: "efs:data:outputs/run-1", MaxRuntime: 7200, + MaxRunAttempts: 3, PathMappings: []lit.JobPathMapping{{ ContainerPath: "/data", ConnectionName: "dataset", @@ -364,6 +398,8 @@ func TestJobRunMapsAdvancedV2Options(t *testing.T) { if got, want := created.ArtifactPath(), "/teamspace/efs_connections/data/outputs/run-1"; got != want { assert.Fail(t, fmt.Sprintf("created artifact path = %q, want %q", got, want)) } + assert.Falsef(t, created.MaxRunAttempts() != 3 || created.CurrentRunAttempt() != 1, + "unexpected run-attempt fields: %d %d", created.MaxRunAttempts(), created.CurrentRunAttempt()) } func TestJobRunMapsScratchDisksForStudioJobs(t *testing.T) { diff --git a/python/lightning_sdk/api/job_api.py b/python/lightning_sdk/api/job_api.py index cfdf0928a..cb8d3b11e 100644 --- a/python/lightning_sdk/api/job_api.py +++ b/python/lightning_sdk/api/job_api.py @@ -4,7 +4,7 @@ from contextlib import suppress from datetime import datetime, timezone from functools import lru_cache -from typing import TYPE_CHECKING, Dict, Iterator, List, Optional, Union +from typing import TYPE_CHECKING, Any, Dict, Iterator, List, Optional, Union from urllib.error import URLError from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse from urllib.request import urlopen @@ -147,6 +147,7 @@ def submit_job( entrypoint: Optional[str], path_mappings: Optional[Dict[str, str]], max_runtime: Optional[int] = None, + max_run_attempts: Optional[int] = None, reuse_snapshot: bool = True, scratch_disks: Optional[Dict[str, int]] = None, placement_group_id: Optional[str] = None, @@ -171,6 +172,8 @@ def submit_job( max_runtime: DWS (Dynamic Workload Scheduler) reservation duration in seconds (e.g. some top-end GCP GPUs). Has no effect on non-DWS or interruptible (spot) machines. ``None`` means no reservation is requested. + max_run_attempts: Max number of run attempts for this job. ``None`` or ``0`` means + unset (backend default). ``1`` means a single attempt (no retries). reuse_snapshot: Whether to reuse the Studio's existing filesystem snapshot. scratch_disks: Optional mapping of scratch-disk mount paths to their sizes in GiB. placement_group_id: Optional placement group identifier for colocating the job. @@ -203,6 +206,7 @@ def submit_job( entrypoint=entrypoint, path_mappings=path_mappings, max_runtime=max_runtime, + max_run_attempts=max_run_attempts, reuse_snapshot=reuse_snapshot, scratch_disks=sanitized_scratch_disks, placement_group_id=placement_group_id, @@ -227,6 +231,7 @@ def _create_job_body( path_mappings: Optional[Dict[str, str]], reuse_snapshot: bool, max_runtime: Optional[int] = None, + max_run_attempts: Optional[int] = None, machine_image_version: Optional[str] = None, scratch_disks: Optional[Dict[str, int]] = None, placement_group_id: Optional[str] = None, @@ -250,6 +255,8 @@ def _create_job_body( max_runtime: DWS (Dynamic Workload Scheduler) reservation duration in seconds (e.g. some top-end GCP GPUs). Has no effect on non-DWS or interruptible (spot) machines. ``None`` means no reservation is requested. + max_run_attempts: Max number of run attempts for this job. ``None`` or ``0`` means + unset (backend default). ``1`` means a single attempt (no retries). machine_image_version: Pinned machine-image version string, or ``None`` for the default. scratch_disks: Optional mapping of scratch-disk mount paths to their sizes in GiB. placement_group_id: Optional placement group identifier for colocating the job. @@ -269,9 +276,11 @@ def _create_job_body( path_mappings_list = resolve_path_mappings(mappings=path_mappings or {}) # need to go via kwargs for typing compatibility since autogenerated apis accept None but aren't typed with None - optional_spec_kwargs = {} + optional_spec_kwargs: Dict[str, Any] = {} if max_runtime: optional_spec_kwargs["requested_run_duration_seconds"] = str(max_runtime) + if max_run_attempts: + optional_spec_kwargs["max_run_attempts"] = max_run_attempts # don't do default dicts, as they'll be mutable. Create a fresh one here scratch_disks = scratch_disks or {} diff --git a/python/lightning_sdk/job.py b/python/lightning_sdk/job.py index 314f4a9f8..04b3a445f 100644 --- a/python/lightning_sdk/job.py +++ b/python/lightning_sdk/job.py @@ -224,6 +224,7 @@ def run( entrypoint: Optional[str] = None, path_mappings: Optional[Dict[str, str]] = None, max_runtime: Optional[int] = None, + max_run_attempts: Optional[int] = None, reuse_snapshot: bool = True, scratch_disks: Optional[Dict[str, int]] = None, placement_group_id: Optional[str] = None, @@ -263,6 +264,9 @@ def run( max_runtime: DWS (Dynamic Workload Scheduler) reservation duration in seconds (e.g. some top-end GCP GPUs). Has no effect on non-DWS or interruptible (spot) machines. ``None`` means no reservation is requested. + max_run_attempts: Max number of run attempts for this job. ``None`` or ``0`` means + unset (backend default). ``1`` means a single attempt (no retries). + ``N > 1`` allows up to ``N`` attempts. reuse_snapshot: Whether to reuse a Studio snapshot when multiple jobs for the same Studio are submitted. Turning this off may result in longer startup times. Defaults to True. scratch_disks: Optional mapping of scratch-disk mount paths to their sizes in GiB. @@ -288,6 +292,8 @@ def run( raise ValueError("A job needs to run on at least one machine") if num_machines > 1 and scratch_disks: raise ValueError("scratch_disks are not supported for multi-machine jobs") + if num_machines > 1 and max_run_attempts: + raise ValueError("max_run_attempts is not supported for multi-machine jobs") if image is None: if not isinstance(studio, Studio): @@ -365,6 +371,7 @@ def run( entrypoint=entrypoint, path_mappings=path_mappings, max_runtime=max_runtime, + max_run_attempts=max_run_attempts, reuse_snapshot=reuse_snapshot, scratch_disks=scratch_disks, placement_group_id=placement_group_id, @@ -388,6 +395,7 @@ def _submit( entrypoint: Optional[str] = None, path_mappings: Optional[Dict[str, str]] = None, max_runtime: Optional[int] = None, + max_run_attempts: Optional[int] = None, reuse_snapshot: bool = True, scratch_disks: Optional[Dict[str, int]] = None, placement_group_id: Optional[str] = None, @@ -397,6 +405,8 @@ def _submit( raise ValueError("A job needs to run on at least one machine") if num_machines > 1 and scratch_disks: raise ValueError("scratch_disks are not supported for multi-machine jobs") + if num_machines > 1 and max_run_attempts: + raise ValueError("max_run_attempts is not supported for multi-machine jobs") if studio is not None: studio_id = studio._studio.id @@ -450,6 +460,9 @@ def _submit( raise ValueError("scratch_disk path cannot contain '..'") self._num_machines = num_machines + extra_submit_kwargs: Dict[str, Any] = {} + if num_machines <= 1: + extra_submit_kwargs["max_run_attempts"] = max_run_attempts submitted = self._job_api.submit_job( name=self.name, command=command, @@ -469,6 +482,7 @@ def _submit( placement_group_id=placement_group_id, num_machines=num_machines, scratch_disks=scratch_disks, + **extra_submit_kwargs, ) if num_machines <= 1 and submitted.name != self._name: warnings.warn( @@ -590,6 +604,20 @@ def private_ip_address(self) -> Optional[str]: def placement_group_id(self) -> Optional[str]: return self._guaranteed_job.spec.placement_group_id + @property + def max_run_attempts(self) -> Optional[int]: + """Max number of run attempts for this job, or ``None`` if unset.""" + spec = getattr(self._guaranteed_job, "spec", None) + value = getattr(spec, "max_run_attempts", None) + return value or None + + @property + def current_run_attempt(self) -> Optional[int]: + """Current run attempt for this job, or ``None`` if unset.""" + spec = getattr(self._guaranteed_job, "spec", None) + value = getattr(spec, "current_run_attempt", None) + return value or None + @property def rank(self) -> Optional[int]: if self.is_multi_machine: diff --git a/python/lightning_sdk/pipeline/steps.py b/python/lightning_sdk/pipeline/steps.py index 136dbf373..d9ce76107 100644 --- a/python/lightning_sdk/pipeline/steps.py +++ b/python/lightning_sdk/pipeline/steps.py @@ -205,6 +205,7 @@ def __init__( entrypoint: str = "sh -c", path_mappings: Optional[Dict[str, str]] = None, max_runtime: Optional[int] = None, + max_run_attempts: Optional[int] = None, wait_for: Union[str, List[str], None] = DEFAULT, reuse_snapshot: bool = True, scratch_disks: Optional[Dict[str, int]] = None, @@ -230,6 +231,9 @@ def __init__( entrypoint: Container entrypoint. Defaults to ``sh -c``. path_mappings: Mappings from container paths to data-connection paths. max_runtime: Maximum runtime in seconds. + max_run_attempts: Max number of run attempts for this job. ``None`` or ``0`` means + unset (backend default). ``1`` means a single attempt (no retries). + ``N > 1`` allows up to ``N`` attempts. Only supported for single-machine jobs. wait_for: Names of steps that must complete before this step starts. reuse_snapshot: Whether to reuse a studio snapshot across jobs. Defaults to True. scratch_disks: Extra volumes to mount under ``/teamspace/scratch``. @@ -239,6 +243,8 @@ def __init__( """ if num_machines < 1: raise ValueError("A job needs to run on at least one machine") + if num_machines > 1 and max_run_attempts: + raise ValueError("max_run_attempts is not supported for multi-machine jobs") self.name = name self.machine = machine or Machine.CPU self.command = command @@ -256,6 +262,7 @@ def __init__( self.entrypoint = entrypoint self.path_mappings = path_mappings self.max_runtime = max_runtime + self.max_run_attempts = max_run_attempts self.wait_for = wait_for self.reuse_snapshot = reuse_snapshot self.scratch_disks = scratch_disks @@ -297,6 +304,8 @@ def to_proto( if self.num_machines > 1: if self.scratch_disks: raise ValueError("scratch_disks are not supported for multi-machine jobs") + if self.max_run_attempts: + raise ValueError("max_run_attempts is not supported for multi-machine jobs") body = MMTApiV2._create_mmt_body( name=cast(str, self.name), num_machines=self.num_machines, @@ -337,6 +346,7 @@ def to_proto( entrypoint=self.entrypoint, path_mappings=self.path_mappings, max_runtime=self.max_runtime, + max_run_attempts=self.max_run_attempts, machine_image_version=machine_image_version, reuse_snapshot=self.reuse_snapshot, scratch_disks=self.scratch_disks, diff --git a/python/lightning_sdk/studio.py b/python/lightning_sdk/studio.py index 55b890c1c..c43e4dc4a 100644 --- a/python/lightning_sdk/studio.py +++ b/python/lightning_sdk/studio.py @@ -781,6 +781,7 @@ def run_job( env: Optional[Dict[str, str]] = None, interruptible: bool = False, reuse_snapshot: bool = True, + max_run_attempts: Optional[int] = None, ) -> "Job": """Run async workloads using the compute environment from your studio. @@ -792,6 +793,8 @@ def run_job( interruptible: Whether the job should run on interruptible instances. They are cheaper but can be preempted. reuse_snapshot: Whether the job should reuse a Studio snapshot when multiple jobs for the same Studio are submitted. Turning this off may result in longer job startup times. Defaults to True. + max_run_attempts: Max number of run attempts for this job. ``None`` or ``0`` means + unset (backend default). ``1`` means a single attempt (no retries). Returns: Job: The submitted :class:`Job` instance. @@ -809,6 +812,7 @@ def run_job( env=env, interruptible=interruptible, reuse_snapshot=reuse_snapshot, + max_run_attempts=max_run_attempts, ) def run_mmt( diff --git a/python/tests/api/test_job_api.py b/python/tests/api/test_job_api.py index 3f529ea21..c829384d0 100644 --- a/python/tests/api/test_job_api.py +++ b/python/tests/api/test_job_api.py @@ -133,6 +133,57 @@ def test_job_v2_submit_job_threads_placement_group_id(mocker_auth): assert body.spec.placement_group_id == "pg-1" +def test_job_v2_submit_job_threads_max_run_attempts(mocker_auth): + job_api = JobApiV2() + create_job_mock = mock.MagicMock() + job_api._client.jobs_service_create_job = create_job_mock + + job_api.submit_job( + name="test-job", + cloud_account="c-abc", + teamspace_id="ts-abc", + image="image-abc", + studio_id="", + machine=Machine.CPU, + interruptible=False, + env=None, + command="echo hello", + image_credentials=None, + cloud_account_auth=False, + entrypoint="sh -c", + path_mappings=None, + max_run_attempts=3, + ) + + body = create_job_mock.call_args.kwargs["body"] + assert body.spec.max_run_attempts == 3 + + +def test_job_v2_submit_job_omits_unset_max_run_attempts(mocker_auth): + job_api = JobApiV2() + create_job_mock = mock.MagicMock() + job_api._client.jobs_service_create_job = create_job_mock + + job_api.submit_job( + name="test-job", + cloud_account="c-abc", + teamspace_id="ts-abc", + image="image-abc", + studio_id="", + machine=Machine.CPU, + interruptible=False, + env=None, + command="echo hello", + image_credentials=None, + cloud_account_auth=False, + entrypoint="sh -c", + path_mappings=None, + ) + + body = create_job_mock.call_args.kwargs["body"] + assert body.spec.max_run_attempts is None + + def test_get_job_by_name(mocker_auth): job_api = JobApiV2() diff --git a/python/tests/core/test_job.py b/python/tests/core/test_job.py index d9dd68abd..bf6d21688 100644 --- a/python/tests/core/test_job.py +++ b/python/tests/core/test_job.py @@ -158,6 +158,7 @@ def test_submit_job_v2_image(internal_studio_init_mocker, machine, command, env, entrypoint=None, path_mappings=None, max_runtime=None, + max_run_attempts=None, reuse_snapshot=True, scratch_disks=None, placement_group_id=None, @@ -258,6 +259,27 @@ def test_job_exposes_private_provisioning_metadata(internal_studio_init_mocker): assert job.rank == 3 +@mock.patch("lightning_sdk.lightning_cloud.rest_client.Auth", new=mock.MagicMock()) +def test_job_exposes_run_attempt_metadata(internal_studio_init_mocker): + teamspace = Teamspace("ts-abc", org="org-abc") + job = Job("test-job", teamspace, _fetch_job=False) + job._job = V1Job( + id="job-123", + name="test-job", + spec=V1JobSpec(max_run_attempts=3, current_run_attempt=2, parent_job_id="job-parent"), + ) + + assert job.max_run_attempts == 3 + assert job.current_run_attempt == 2 + assert job._guaranteed_job.spec.parent_job_id == "job-parent" + + unset = Job("unset-job", teamspace, _fetch_job=False) + unset._job = V1Job(id="job-456", name="unset-job", spec=V1JobSpec()) + assert unset.max_run_attempts is None + assert unset.current_run_attempt is None + assert unset._guaranteed_job.spec.parent_job_id is None + + @mock.patch("lightning_sdk.lightning_cloud.rest_client.Auth", new=mock.MagicMock()) def test_job_exposes_start_and_stop_times(internal_studio_init_mocker): teamspace = Teamspace("ts-abc", org="org-abc") @@ -326,6 +348,40 @@ def test_submit_job_threads_placement_group_id(internal_studio_init_mocker): assert submit_mock.call_args.kwargs["placement_group_id"] == "pg-1" +@mock.patch("lightning_sdk.lightning_cloud.rest_client.Auth", new=mock.MagicMock()) +def test_submit_job_threads_max_run_attempts(internal_studio_init_mocker): + teamspace = Teamspace("ts-abc", org="org-abc") + job = Job("test-job", teamspace, _fetch_job=False) + submit_mock = mock.MagicMock() + job._job_api.submit_job = submit_mock + + job._submit( + machine=Machine.CPU, + image="image-abc", + command="echo hello", + cloud_account="c-abc", + max_run_attempts=3, + ) + + assert submit_mock.call_args.kwargs["max_run_attempts"] == 3 + + +@mock.patch("lightning_sdk.lightning_cloud.rest_client.Auth", new=mock.MagicMock()) +def test_submit_rejects_max_run_attempts_for_multi_machine(internal_studio_init_mocker): + teamspace = Teamspace("ts-abc", org="org-abc") + job = Job("test-job", teamspace, _fetch_job=False) + + with pytest.raises(ValueError, match="max_run_attempts is not supported for multi-machine jobs"): + job._submit( + machine=Machine.CPU, + image="image-abc", + command="echo hello", + cloud_account="c-abc", + num_machines=2, + max_run_attempts=3, + ) + + @pytest.mark.parametrize("machine", [Machine.L4, Machine.DATA_PREP_MAX]) @pytest.mark.parametrize("env", [None, {"key": "value"}]) @pytest.mark.parametrize("interruptible", [True, False]) @@ -360,6 +416,7 @@ def test_submit_job_v2_studio(internal_studio_init_mocker, machine, env, interru entrypoint=None, path_mappings=None, max_runtime=None, + max_run_attempts=None, reuse_snapshot=True, scratch_disks=None, placement_group_id=None, @@ -748,6 +805,7 @@ def test_submit_jobv2_studio_resolve( entrypoint=None, path_mappings=None, max_runtime=None, + max_run_attempts=None, reuse_snapshot=True, scratch_disks=None, placement_group_id=None, @@ -846,6 +904,7 @@ def test_submit_job_v2_image_from_studio( entrypoint="sh -c", path_mappings=None, max_runtime=None, + max_run_attempts=None, reuse_snapshot=True, scratch_disks=None, placement_group_id=None, @@ -891,6 +950,7 @@ def test_run_job_with_cloud_provider( entrypoint="sh -c", path_mappings=None, max_runtime=None, + max_run_attempts=None, reuse_snapshot=True, scratch_disks=None, placement_group_id=None, diff --git a/python/tests/core/test_pipeline.py b/python/tests/core/test_pipeline.py index c096bae9e..18fc4c5e2 100644 --- a/python/tests/core/test_pipeline.py +++ b/python/tests/core/test_pipeline.py @@ -173,6 +173,19 @@ def test_job_step_threads_placement_group_id(): assert proto.job.spec.placement_group_id == "pg-1" +@patch("lightning_sdk.pipeline.steps.CloudAccountApi", new=MagicMock()) +def test_job_step_threads_max_run_attempts(): + job = JobStep(name="job-0", machine=Machine.CPU, max_run_attempts=3) + proto = job.to_proto(MagicMock(), "", False) + + assert proto.job.spec.max_run_attempts == 3 + + +def test_job_step_rejects_max_run_attempts_for_multi_machine(): + with pytest.raises(ValueError, match="max_run_attempts is not supported for multi-machine jobs"): + JobStep(name="job-0", machine=Machine.CPU, num_machines=3, max_run_attempts=3) + + @patch("lightning_sdk.pipeline.steps.CloudAccountApi", new=MagicMock()) def test_job_step_supports_multiple_machines(): job = JobStep(name="job-0", machine=Machine.CPU, num_machines=3)