Skip to content
Open
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
1 change: 1 addition & 0 deletions .nextchanges/cli/sandbox-transient-stopped.md
Original file line number Diff line number Diff line change
@@ -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))
15 changes: 13 additions & 2 deletions cmd/sandbox/ssh_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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,
},
} {
Expand Down
23 changes: 17 additions & 6 deletions cmd/sandbox/start.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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 {
Expand All @@ -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()
Expand Down
254 changes: 254 additions & 0 deletions cmd/sandbox/start_test.go
Original file line number Diff line number Diff line change
@@ -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)
})
})
}
}