diff --git a/.nextchanges/cli/sandbox-transient-stopped.md b/.nextchanges/cli/sandbox-transient-stopped.md new file mode 100644 index 00000000000..786918fa120 --- /dev/null +++ b/.nextchanges/cli/sandbox-transient-stopped.md @@ -0,0 +1 @@ +* Keep waiting through a transient `Stopped` status during `databricks sandbox start` and `databricks sandbox ssh`, retrying start requests at a throttled interval until the sandbox is running or the wait times out. ([#7010](https://github.com/databricks/cli/pull/7010)) diff --git a/cmd/sandbox/ssh_test.go b/cmd/sandbox/ssh_test.go index 31f53eee6b4..5cc440ee01f 100644 --- a/cmd/sandbox/ssh_test.go +++ b/cmd/sandbox/ssh_test.go @@ -55,13 +55,24 @@ func TestEnsureRunningLifecycle(t *testing.T) { wantStatus: "Running", }, { - name: "stopped after start remains an error", + name: "transient stopped after start is tolerated", initialStatus: "Stopped", steps: []requestStep{ {method: http.MethodPost, path: sandboxPath("test-id") + "/start", status: "Pending"}, {method: http.MethodGet, path: sandboxPath("test-id"), status: "Stopped"}, + {method: http.MethodGet, path: sandboxPath("test-id"), status: "Running"}, + }, + wantStatus: "Running", + wantStarts: 1, + }, + { + name: "failed after start remains an error", + initialStatus: "Stopped", + steps: []requestStep{ + {method: http.MethodPost, path: sandboxPath("test-id") + "/start", status: "Pending"}, + {method: http.MethodGet, path: sandboxPath("test-id"), status: "Failed"}, }, - wantErr: `sandbox test-id reached unexpected state "Stopped" while starting`, + wantErr: `sandbox test-id reached unexpected state "Failed" while starting`, wantStarts: 1, }, } { diff --git a/cmd/sandbox/start.go b/cmd/sandbox/start.go index ba48a42efd9..9be871d0f02 100644 --- a/cmd/sandbox/start.go +++ b/cmd/sandbox/start.go @@ -17,8 +17,9 @@ import ( // reaches Running. 10 min covers the observed cold-start range // (5–13 min); stuck sandboxes surface as a timeout, not a hang. const ( - startPollInterval = 2 * time.Second - startWaitTimeout = 10 * time.Minute + startPollInterval = 2 * time.Second + startWaitTimeout = 10 * time.Minute + startRenudgeInterval = 15 * time.Second ) func newStartCommand() *cobra.Command { @@ -93,21 +94,27 @@ Example: return cmd } -func waitForRunning(ctx context.Context, api *sandboxAPI, s *spinner, id string) (*sandboxEntry, error) { - return waitForState(ctx, api, s, id, "Running", "Starting", "Stopped", "Terminated", "Failed") +type runWaiter interface { + get(ctx context.Context, id string) (*sandboxEntry, error) + start(ctx context.Context, id string) (*sandboxEntry, error) +} + +func waitForRunning(ctx context.Context, api runWaiter, s *spinner, id string) (*sandboxEntry, error) { + return waitForState(ctx, api, s, id, "Running", "Starting", "Terminated", "Failed") } // The API rejects start requests until teardown finishes, so STOPPING needs // its own polling phase. -func waitForStopped(ctx context.Context, api *sandboxAPI, s *spinner, id string) (*sandboxEntry, error) { +func waitForStopped(ctx context.Context, api runWaiter, s *spinner, id string) (*sandboxEntry, error) { return waitForState(ctx, api, s, id, "Stopped", "Stopping", "Terminated", "Failed") } // Centralizing lifecycle polling keeps timeout, cancellation, and terminal // state handling consistent across start and SSH. -func waitForState(ctx context.Context, api *sandboxAPI, s *spinner, id, targetStatus, operation string, unexpectedStatuses ...string) (*sandboxEntry, error) { +func waitForState(ctx context.Context, api runWaiter, s *spinner, id, targetStatus, operation string, unexpectedStatuses ...string) (*sandboxEntry, error) { start := time.Now() deadline := start.Add(startWaitTimeout) + lastStart := start for { sb, err := api.get(ctx, id) if err != nil { @@ -126,6 +133,10 @@ func waitForState(ctx context.Context, api *sandboxAPI, s *spinner, id, targetSt if time.Now().After(deadline) { return nil, fmt.Errorf("sandbox %s did not reach %s within %s (last seen %s)", id, targetStatus, startWaitTimeout, sb.Status) } + if targetStatus == "Running" && strings.EqualFold(sb.Status, "stopped") && time.Since(lastStart) >= startRenudgeInterval { + _, _ = api.start(ctx, id) + lastStart = time.Now() + } select { case <-ctx.Done(): return nil, ctx.Err() diff --git a/cmd/sandbox/start_test.go b/cmd/sandbox/start_test.go new file mode 100644 index 00000000000..97a484a649e --- /dev/null +++ b/cmd/sandbox/start_test.go @@ -0,0 +1,254 @@ +package sandbox + +import ( + "context" + "errors" + "slices" + "testing" + "testing/synctest" + "time" + + "github.com/databricks/cli/libs/cmdio" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +const waitTestSandboxID = "test-id" + +type fakeRunWaiter struct { + statuses []string + getErr error + startErr error + polls int + starts []time.Time +} + +func (waiter *fakeRunWaiter) get(_ context.Context, id string) (*sandboxEntry, error) { + if waiter.getErr != nil { + return nil, waiter.getErr + } + status := waiter.statuses[min(waiter.polls, len(waiter.statuses)-1)] + waiter.polls++ + return &sandboxEntry{SandboxID: id, Status: status}, nil +} + +func (waiter *fakeRunWaiter) start(_ context.Context, id string) (*sandboxEntry, error) { + waiter.starts = append(waiter.starts, time.Now()) + if waiter.startErr != nil { + return nil, waiter.startErr + } + return &sandboxEntry{SandboxID: id, Status: "Creating"}, nil +} + +func TestWaitForRunning(t *testing.T) { + t.Parallel() + + for _, testCase := range []struct { + name string + statuses []string + startErr error + wantErr string + wantStarts int + }{ + { + name: "already running", + statuses: []string{"Running"}, + }, + { + name: "creating", + statuses: []string{"Creating", "Creating", "Running"}, + }, + { + name: "transient stopped", + statuses: []string{"Stopped", "Creating", "Running"}, + }, + { + name: "stopped retries start", + statuses: append(slices.Repeat([]string{"Stopped"}, 9), "Creating", "Running"), + wantStarts: 1, + }, + { + name: "case insensitive statuses", + statuses: append(slices.Repeat([]string{"sToPpEd"}, 9), "rUnNiNg"), + wantStarts: 1, + }, + { + name: "retries are throttled", + statuses: append(slices.Repeat([]string{"Stopped"}, 18), "Running"), + wantStarts: 2, + }, + { + name: "retry error is tolerated", + statuses: append(slices.Repeat([]string{"Stopped"}, 9), "Creating", "Running"), + startErr: errors.New("currently stopping"), + wantStarts: 1, + }, + { + name: "failed is terminal", + statuses: []string{"Failed"}, + wantErr: `sandbox test-id reached unexpected state "Failed" while starting`, + }, + { + name: "terminated is terminal", + statuses: []string{"Terminated"}, + wantErr: `sandbox test-id reached unexpected state "Terminated" while starting`, + }, + } { + t.Run(testCase.name, func(t *testing.T) { + t.Parallel() + synctest.Test(t, func(t *testing.T) { + ctx := cmdio.MockDiscard(t.Context()) + progress := spin(ctx, "") + defer progress.Close() + waiter := &fakeRunWaiter{statuses: testCase.statuses, startErr: testCase.startErr} + started := time.Now() + + got, err := waitForRunning(ctx, waiter, progress, waitTestSandboxID) + if testCase.wantErr != "" { + require.EqualError(t, err, testCase.wantErr) + assert.Nil(t, got) + } else { + require.NoError(t, err) + assert.Equal(t, waitTestSandboxID, got.SandboxID) + assert.Equal(t, testCase.statuses[len(testCase.statuses)-1], got.Status) + } + assert.Equal(t, len(testCase.statuses), waiter.polls) + assert.Len(t, waiter.starts, testCase.wantStarts) + lastStart := started + for _, retried := range waiter.starts { + assert.GreaterOrEqual(t, retried.Sub(lastStart), startRenudgeInterval) + lastStart = retried + } + }) + }) + } +} + +func TestWaitForRunningTimeout(t *testing.T) { + t.Parallel() + + for _, status := range []string{"Creating", "Stopped"} { + t.Run(status, func(t *testing.T) { + t.Parallel() + synctest.Test(t, func(t *testing.T) { + ctx := cmdio.MockDiscard(t.Context()) + progress := spin(ctx, "") + defer progress.Close() + waiter := &fakeRunWaiter{statuses: []string{status}} + started := time.Now() + + got, err := waitForRunning(ctx, waiter, progress, waitTestSandboxID) + require.ErrorContains(t, err, "did not reach Running within") + assert.ErrorContains(t, err, "last seen "+status) + assert.Nil(t, got) + assert.GreaterOrEqual(t, time.Since(started), startWaitTimeout) + assert.LessOrEqual(t, time.Since(started), startWaitTimeout+startPollInterval) + if status == "Stopped" { + require.NotEmpty(t, waiter.starts) + lastStart := started + for _, retried := range waiter.starts { + assert.GreaterOrEqual(t, retried.Sub(lastStart), startRenudgeInterval) + assert.LessOrEqual(t, retried.Sub(started), startWaitTimeout) + lastStart = retried + } + } else { + assert.Empty(t, waiter.starts) + } + }) + }) + } +} + +func TestWaitForRunningPollingError(t *testing.T) { + t.Parallel() + + ctx := cmdio.MockDiscard(t.Context()) + progress := spin(ctx, "") + defer progress.Close() + pollErr := errors.New("sandbox not found") + waiter := &fakeRunWaiter{getErr: pollErr} + + got, err := waitForRunning(ctx, waiter, progress, waitTestSandboxID) + require.ErrorIs(t, err, pollErr) + assert.ErrorContains(t, err, "polling status of "+waitTestSandboxID) + assert.Nil(t, got) + assert.Empty(t, waiter.starts) +} + +func TestWaitForRunningCancellation(t *testing.T) { + t.Parallel() + + for _, cancelAfter := range []time.Duration{0, 5 * time.Second, 17 * time.Second} { + t.Run(cancelAfter.String(), func(t *testing.T) { + t.Parallel() + synctest.Test(t, func(t *testing.T) { + ctx, cancel := context.WithCancel(cmdio.MockDiscard(t.Context())) + defer cancel() + if cancelAfter == 0 { + cancel() + } else { + timer := time.AfterFunc(cancelAfter, cancel) + defer timer.Stop() + } + progress := spin(ctx, "") + defer progress.Close() + waiter := &fakeRunWaiter{statuses: []string{"Stopped"}} + started := time.Now() + + got, err := waitForRunning(ctx, waiter, progress, waitTestSandboxID) + assert.ErrorIs(t, err, context.Canceled) + assert.Nil(t, got) + assert.Equal(t, cancelAfter, time.Since(started)) + }) + }) + } +} + +func TestWaitForStoppedDoesNotRestart(t *testing.T) { + t.Parallel() + + for _, testCase := range []struct { + name string + statuses []string + wantErr string + }{ + { + name: "already stopped", + statuses: []string{"Stopped"}, + }, + { + name: "stopping", + statuses: []string{"Stopping", "Stopped"}, + }, + { + name: "failed is terminal", + statuses: []string{"Failed"}, + wantErr: `sandbox test-id reached unexpected state "Failed" while stopping`, + }, + { + name: "terminated is terminal", + statuses: []string{"Terminated"}, + wantErr: `sandbox test-id reached unexpected state "Terminated" while stopping`, + }, + } { + t.Run(testCase.name, func(t *testing.T) { + t.Parallel() + synctest.Test(t, func(t *testing.T) { + ctx := cmdio.MockDiscard(t.Context()) + progress := spin(ctx, "") + defer progress.Close() + waiter := &fakeRunWaiter{statuses: testCase.statuses} + + got, err := waitForStopped(ctx, waiter, progress, waitTestSandboxID) + if testCase.wantErr != "" { + require.EqualError(t, err, testCase.wantErr) + assert.Nil(t, got) + } else { + require.NoError(t, err) + assert.Equal(t, "Stopped", got.Status) + } + assert.Empty(t, waiter.starts) + }) + }) + } +}