diff --git a/.nextchanges/cli/ssh-setup-server-lifecycle-flags.md b/.nextchanges/cli/ssh-setup-server-lifecycle-flags.md new file mode 100644 index 00000000000..f0712f83810 --- /dev/null +++ b/.nextchanges/cli/ssh-setup-server-lifecycle-flags.md @@ -0,0 +1 @@ +* Add `--max-clients` and `--server-timeout` flags to `databricks ssh setup`, and `--server-timeout` to `databricks ssh connect`. Both are fixed when the SSH tunnel server job is submitted, so `ssh setup` now serializes them into the generated `ProxyCommand` instead of falling back to the built-in defaults. ([#6547](https://github.com/databricks/cli/pull/6547)) diff --git a/acceptance/ssh/setup/out.test.toml b/acceptance/ssh/setup/out.test.toml new file mode 100644 index 00000000000..0938e678987 --- /dev/null +++ b/acceptance/ssh/setup/out.test.toml @@ -0,0 +1,2 @@ +Cloud = false +EnvMatrix.DATABRICKS_BUNDLE_ENGINE = ["direct"] diff --git a/acceptance/ssh/setup/output.txt b/acceptance/ssh/setup/output.txt new file mode 100644 index 00000000000..589680305ce --- /dev/null +++ b/acceptance/ssh/setup/output.txt @@ -0,0 +1,18 @@ + +=== ProxyCommand written by setup --max-clients=25 --server-timeout=48h +ssh connect --proxy --cluster=[TEST_DEFAULT_CLUSTER_ID] --auto-start-cluster=true --shutdown-delay=10m0s --max-clients=25 --server-timeout=48h0m0s + +=== A shutdown delay beyond the default lifetime raises it, no --server-timeout needed +ssh connect --proxy --cluster=[TEST_DEFAULT_CLUSTER_ID] --auto-start-cluster=true --shutdown-delay=48h0m0s --max-clients=10 --server-timeout=48h0m0s + +=== Rejects a server that would refuse every connection +>>> [CLI] ssh setup --name=broken --cluster=[TEST_DEFAULT_CLUSTER_ID] --max-clients=0 +Error: --max-clients must be at least 1, got 0 + +=== Rejects a shutdown delay the server can never reach +>>> [CLI] ssh setup --name=broken --cluster=[TEST_DEFAULT_CLUSTER_ID] --shutdown-delay=48h --server-timeout=24h +Error: --shutdown-delay (48h0m0s) cannot be longer than --server-timeout (24h0m0s) + +=== No host config is written for the rejected setups +home/.databricks/ssh-tunnel-configs/long-delay +home/.databricks/ssh-tunnel-configs/my-cluster diff --git a/acceptance/ssh/setup/script b/acceptance/ssh/setup/script new file mode 100644 index 00000000000..a406a2e4c80 --- /dev/null +++ b/acceptance/ssh/setup/script @@ -0,0 +1,27 @@ +sethome "./home" + +# Send the command's own output to a LOG file: it echoes the whole host config, whose IdentityFile +# and ProxyCommand carry OS-specific absolute paths. +$CLI ssh setup --name=my-cluster --cluster=$TEST_DEFAULT_CLUSTER_ID --max-clients=25 --server-timeout=48h &>LOG.setup + +# The ProxyCommand is the invocation that submits the SSH server job, so --max-clients and +# --server-timeout have to appear there or `ssh my-cluster` silently starts a server with the +# defaults. +title "ProxyCommand written by setup --max-clients=25 --server-timeout=48h\n" +sed -n 's/.*\(ssh connect --proxy.*\)/\1/p' "$HOME/.databricks/ssh-tunnel-configs/my-cluster" + +# A shutdown delay longer than the default 24h lifetime worked before --server-timeout existed, +# and `ssh setup` persisted it into the ProxyCommand. OpenSSH runs that ProxyCommand verbatim, so +# the delay has to keep raising the lifetime rather than becoming an error the user cannot fix. +title "A shutdown delay beyond the default lifetime raises it, no --server-timeout needed\n" +$CLI ssh setup --name=long-delay --cluster=$TEST_DEFAULT_CLUSTER_ID --shutdown-delay=48h &>LOG.long-delay +sed -n 's/.*\(ssh connect --proxy.*\)/\1/p' "$HOME/.databricks/ssh-tunnel-configs/long-delay" + +title "Rejects a server that would refuse every connection" +musterr trace $CLI ssh setup --name=broken --cluster=$TEST_DEFAULT_CLUSTER_ID --max-clients=0 + +title "Rejects a shutdown delay the server can never reach" +musterr trace $CLI ssh setup --name=broken --cluster=$TEST_DEFAULT_CLUSTER_ID --shutdown-delay=48h --server-timeout=24h + +title "No host config is written for the rejected setups\n" +find.py 'ssh-tunnel-configs' --expect 2 diff --git a/acceptance/ssh/setup/test.toml b/acceptance/ssh/setup/test.toml new file mode 100644 index 00000000000..5970721ec19 --- /dev/null +++ b/acceptance/ssh/setup/test.toml @@ -0,0 +1,8 @@ +# `ssh setup` only reads the cluster and writes local SSH config, so a real workspace adds nothing. +Cloud = false + +Ignore = [ + "home", +] + +EnvMatrix.DATABRICKS_BUNDLE_ENGINE = ["direct"] diff --git a/experimental/ssh/cmd/connect.go b/experimental/ssh/cmd/connect.go index 7502d70f560..adf12e46db6 100644 --- a/experimental/ssh/cmd/connect.go +++ b/experimental/ssh/cmd/connect.go @@ -7,8 +7,25 @@ import ( "github.com/databricks/cli/experimental/ssh/internal/client" "github.com/databricks/cli/libs/cmdctx" "github.com/spf13/cobra" + "github.com/spf13/pflag" ) +// resolveServerTimeout returns the lifetime to submit the SSH tunnel server job with. +// +// Before --server-timeout existed the lifetime was max(24h, --shutdown-delay), so +// `ssh setup --shutdown-delay=48h` produced a working 48h tunnel and persisted that delay into +// the generated ProxyCommand. Keep honoring it whenever --server-timeout is not set: OpenSSH +// runs a persisted ProxyCommand verbatim, so a user whose host config predates the flag cannot +// add it, and rejecting the pair would break `ssh ` with an error naming a flag they have +// no way to pass. An explicit --server-timeout always wins, and ClientOptions.Validate still +// rejects a shutdown delay that outlives a lifetime the user asked for explicitly. +func resolveServerTimeout(flags *pflag.FlagSet, serverTimeout, shutdownDelay time.Duration) time.Duration { + if flags.Changed("server-timeout") { + return serverTimeout + } + return max(serverTimeout, shutdownDelay) +} + func newConnectCommand() *cobra.Command { cmd := &cobra.Command{ Use: "connect", @@ -32,6 +49,7 @@ Connect to a dedicated cluster: var serverMetadata string var shutdownDelay time.Duration var maxClients int + var serverTimeout time.Duration var handoverTimeout time.Duration var releasesDir string var autoStartCluster bool @@ -46,6 +64,7 @@ Connect to a dedicated cluster: cmd.Flags().StringVar(&clusterID, "cluster", "", "Databricks dedicated cluster ID") cmd.Flags().DurationVar(&shutdownDelay, "shutdown-delay", defaultShutdownDelay, "Delay before shutting down the server after the last client disconnects") cmd.Flags().IntVar(&maxClients, "max-clients", defaultMaxClients, "Maximum number of SSH clients") + cmd.Flags().DurationVar(&serverTimeout, "server-timeout", defaultServerTimeout, "Maximum lifetime of the SSH server; it is terminated after this duration even if clients are connected") cmd.Flags().BoolVar(&autoStartCluster, "auto-start-cluster", true, "Automatically start the cluster if it is not running") cmd.Flags().StringVar(&connectionName, "name", "", "Connection name to reuse across sessions (serverless only)") @@ -121,7 +140,7 @@ Connect to a dedicated cluster: HandoverTimeout: handoverTimeout, KeepaliveInterval: defaultKeepaliveInterval, ReleasesDir: releasesDir, - ServerTimeout: max(serverTimeout, shutdownDelay), + ServerTimeout: resolveServerTimeout(cmd.Flags(), serverTimeout, shutdownDelay), TaskStartupTimeout: startupTimeout, AutoStartCluster: autoStartCluster, ClientPublicKeyName: clientPublicKeyName, diff --git a/experimental/ssh/cmd/constants.go b/experimental/ssh/cmd/constants.go index dd5d3b2fdf2..001ee7f23a6 100644 --- a/experimental/ssh/cmd/constants.go +++ b/experimental/ssh/cmd/constants.go @@ -13,8 +13,10 @@ const ( // 30 second SSH-level keepalive that was verified to prevent it. defaultKeepaliveInterval = 20 * time.Second defaultEnvironmentVersion = 4 + // Default cap on how long an SSH tunnel server is allowed to live. Fixed when the + // server job is submitted, so it is only settable by the invocation that starts it. + defaultServerTimeout = 24 * time.Hour - serverTimeout = 24 * time.Hour taskStartupTimeout = 10 * time.Minute gpuTaskStartupTimeout = 45 * time.Minute serverPortRange = 100 diff --git a/experimental/ssh/cmd/setup.go b/experimental/ssh/cmd/setup.go index ff67501445e..a8848e741f8 100644 --- a/experimental/ssh/cmd/setup.go +++ b/experimental/ssh/cmd/setup.go @@ -24,6 +24,8 @@ For serverless connections, use ` + "`databricks ssh connect`" + ` (no setup ste var clusterID string var sshConfigPath string var shutdownDelay time.Duration + var maxClients int + var serverTimeout time.Duration var autoStartCluster bool var autoApprove bool @@ -33,6 +35,8 @@ For serverless connections, use ` + "`databricks ssh connect`" + ` (no setup ste cmd.Flags().BoolVar(&autoStartCluster, "auto-start-cluster", true, "Automatically start the cluster when establishing the ssh connection") cmd.Flags().StringVar(&sshConfigPath, "ssh-config", "", "Path to SSH config file (default ~/.ssh/config)") cmd.Flags().DurationVar(&shutdownDelay, "shutdown-delay", defaultShutdownDelay, "SSH server will terminate after this delay if there are no active connections") + cmd.Flags().IntVar(&maxClients, "max-clients", defaultMaxClients, "Maximum number of SSH clients") + cmd.Flags().DurationVar(&serverTimeout, "server-timeout", defaultServerTimeout, "Maximum lifetime of the SSH server; it is terminated after this duration even if clients are connected") cmd.Flags().BoolVar(&autoApprove, "auto-approve", false, "Skip confirmation prompts, recreating existing SSH host configs without asking") cmd.PreRunE = func(cmd *cobra.Command, args []string) error { @@ -51,6 +55,8 @@ For serverless connections, use ` + "`databricks ssh connect`" + ` (no setup ste AutoStartCluster: autoStartCluster, SSHConfigPath: sshConfigPath, ShutdownDelay: shutdownDelay, + MaxClients: maxClients, + ServerTimeout: resolveServerTimeout(cmd.Flags(), serverTimeout, shutdownDelay), Profile: wsClient.Config.Profile, AutoApprove: autoApprove, } diff --git a/experimental/ssh/cmd/setup_test.go b/experimental/ssh/cmd/setup_test.go new file mode 100644 index 00000000000..9cce1f128c6 --- /dev/null +++ b/experimental/ssh/cmd/setup_test.go @@ -0,0 +1,183 @@ +package ssh + +import ( + "strings" + "testing" + "time" + + "github.com/databricks/cli/experimental/ssh/internal/client" + "github.com/spf13/cobra" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// The ProxyCommand written by `ssh setup` is what OpenSSH later executes, and it is the +// invocation that submits the SSH server job. So every flag setup serializes has to exist on +// `ssh connect` and survive the round trip - otherwise the value the user passed to setup is +// silently replaced by connect's own default. +func TestSetupProxyCommandRoundTripsThroughConnect(t *testing.T) { + opts := client.ClientOptions{ + ClusterID: "cluster-123", + AutoStartCluster: true, + ShutdownDelay: 15 * time.Minute, + MaxClients: 25, + ServerTimeout: 48 * time.Hour, + } + proxyCommand, err := opts.ToProxyCommand() + require.NoError(t, err) + + _, args, found := strings.Cut(proxyCommand, " ssh connect ") + require.True(t, found, "proxy command %q must invoke 'ssh connect'", proxyCommand) + + flags := newConnectCommand().Flags() + require.NoError(t, flags.Parse(strings.Fields(args))) + + proxyMode, err := flags.GetBool("proxy") + require.NoError(t, err) + assert.True(t, proxyMode) + + clusterID, err := flags.GetString("cluster") + require.NoError(t, err) + assert.Equal(t, opts.ClusterID, clusterID) + + autoStart, err := flags.GetBool("auto-start-cluster") + require.NoError(t, err) + assert.Equal(t, opts.AutoStartCluster, autoStart) + + shutdownDelay, err := flags.GetDuration("shutdown-delay") + require.NoError(t, err) + assert.Equal(t, opts.ShutdownDelay, shutdownDelay) + + maxClients, err := flags.GetInt("max-clients") + require.NoError(t, err) + assert.Equal(t, opts.MaxClients, maxClients) + + serverTimeout, err := flags.GetDuration("server-timeout") + require.NoError(t, err) + assert.Equal(t, opts.ServerTimeout, serverTimeout) +} + +// Both commands start the same server, so a value the user does not override has to mean the +// same thing whether the tunnel was configured with setup or started with connect. +func TestSetupAndConnectShareServerLifecycleDefaults(t *testing.T) { + setupFlags := newSetupCommand().Flags() + connectFlags := newConnectCommand().Flags() + + for _, name := range []string{"shutdown-delay", "max-clients", "server-timeout"} { + setupFlag := setupFlags.Lookup(name) + connectFlag := connectFlags.Lookup(name) + require.NotNil(t, setupFlag, "setup is missing --%s", name) + require.NotNil(t, connectFlag, "connect is missing --%s", name) + assert.Equal(t, connectFlag.DefValue, setupFlag.DefValue, "--%s default differs between setup and connect", name) + } +} + +// A ProxyCommand written by `ssh setup --shutdown-delay=48h` before --server-timeout existed +// carries no --server-timeout, and OpenSSH runs it verbatim - so the delay has to keep raising +// the server's lifetime the way it did then, or `ssh ` breaks for a host config the user +// cannot edit from the command line. +func TestLegacyProxyCommandWithLongShutdownDelayStillWorks(t *testing.T) { + flags := newConnectCommand().Flags() + require.NoError(t, flags.Parse([]string{ + "--proxy", "--cluster=abc-123", "--auto-start-cluster=true", "--shutdown-delay=48h0m0s", + })) + + shutdownDelay, err := flags.GetDuration("shutdown-delay") + require.NoError(t, err) + serverTimeout, err := flags.GetDuration("server-timeout") + require.NoError(t, err) + + opts := client.ClientOptions{ + ProxyMode: true, + ClusterID: "abc-123", + MaxClients: defaultMaxClients, + ShutdownDelay: shutdownDelay, + ServerTimeout: resolveServerTimeout(flags, serverTimeout, shutdownDelay), + } + require.NoError(t, opts.Validate()) + assert.Equal(t, 48*time.Hour, opts.ServerTimeout, "the shutdown delay must raise the server's lifetime to cover it") +} + +// The reason the max() above is gated: an explicitly requested lifetime must not be silently +// widened by a longer shutdown delay. That pair is a real conflict, so it stays an error. +func TestExplicitServerTimeoutShorterThanShutdownDelayIsRejected(t *testing.T) { + flags := newConnectCommand().Flags() + require.NoError(t, flags.Parse([]string{ + "--proxy", "--cluster=abc-123", "--shutdown-delay=48h", "--server-timeout=24h", + })) + + shutdownDelay, err := flags.GetDuration("shutdown-delay") + require.NoError(t, err) + serverTimeout, err := flags.GetDuration("server-timeout") + require.NoError(t, err) + + resolved := resolveServerTimeout(flags, serverTimeout, shutdownDelay) + assert.Equal(t, 24*time.Hour, resolved, "an explicit --server-timeout wins over a longer --shutdown-delay") + + opts := client.ClientOptions{ + ProxyMode: true, + ClusterID: "abc-123", + MaxClients: defaultMaxClients, + ShutdownDelay: shutdownDelay, + ServerTimeout: resolved, + } + assert.EqualError(t, opts.Validate(), "--shutdown-delay (48h0m0s) cannot be longer than --server-timeout (24h0m0s)") +} + +func TestResolveServerTimeout(t *testing.T) { + tests := []struct { + name string + args []string + want time.Duration + }{ + { + name: "no flags: the default lifetime", + args: nil, + want: defaultServerTimeout, + }, + { + name: "shutdown delay within the default lifetime leaves it alone", + args: []string{"--shutdown-delay=1h"}, + want: defaultServerTimeout, + }, + { + name: "shutdown delay beyond the default lifetime raises it", + args: []string{"--shutdown-delay=48h"}, + want: 48 * time.Hour, + }, + { + name: "an explicit lifetime is used as given", + args: []string{"--server-timeout=1h"}, + want: time.Hour, + }, + { + name: "an explicit lifetime is not widened by a longer shutdown delay", + args: []string{"--shutdown-delay=2h", "--server-timeout=1h"}, + want: time.Hour, + }, + { + // Same value as the default, but passed explicitly: the user asked for it, so a + // longer shutdown delay must not override it. + name: "an explicit lifetime equal to the default still wins", + args: []string{"--shutdown-delay=48h", "--server-timeout=24h"}, + want: defaultServerTimeout, + }, + } + + // Both commands resolve the lifetime the same way, so check them together. + for _, newCmd := range []func() *cobra.Command{newSetupCommand, newConnectCommand} { + for _, tt := range tests { + t.Run(newCmd().Name()+"/"+tt.name, func(t *testing.T) { + flags := newCmd().Flags() + require.NoError(t, flags.Parse(tt.args)) + + shutdownDelay, err := flags.GetDuration("shutdown-delay") + require.NoError(t, err) + serverTimeout, err := flags.GetDuration("server-timeout") + require.NoError(t, err) + + assert.Equal(t, tt.want, resolveServerTimeout(flags, serverTimeout, shutdownDelay)) + }) + } + } +} diff --git a/experimental/ssh/internal/client/client.go b/experimental/ssh/internal/client/client.go index 031bbd4d7af..4c7347416de 100644 --- a/experimental/ssh/internal/client/client.go +++ b/experimental/ssh/internal/client/client.go @@ -155,6 +155,22 @@ func (o *ClientOptions) Validate() error { if o.BaseEnvironment != "" && o.ClusterID != "" { return errors.New("--base-environment can only be used with serverless compute") } + // A server started with fewer than one client slot rejects every connection with an + // opaque websocket handshake failure, so catch it here instead. + if o.MaxClients < 1 { + return fmt.Errorf("--max-clients must be at least 1, got %d", o.MaxClients) + } + // The submitted job carries this as timeout_seconds, a whole number of seconds, and 0 means + // "no timeout" in the Jobs API. So every value below one second - not just zero - truncates + // to an unbounded run instead of the short-lived server that was asked for. + if o.ServerTimeout < time.Second { + return fmt.Errorf("--server-timeout must be at least 1s, got %s", o.ServerTimeout) + } + // The server only starts counting down the shutdown delay once the last client leaves, so a + // delay longer than the server's lifetime can never elapse. + if o.ShutdownDelay > o.ServerTimeout { + return fmt.Errorf("--shutdown-delay (%s) cannot be longer than --server-timeout (%s)", o.ShutdownDelay, o.ServerTimeout) + } return nil } @@ -229,6 +245,18 @@ func (o *ClientOptions) ToProxyCommand() (string, error) { executablePath, o.ClusterID, o.AutoStartCluster, o.ShutdownDelay.String()) } + // Both of these are fixed when the server job is submitted, and for a host configured by + // `ssh setup` the submitting invocation is always the ProxyCommand, so they have to be + // carried here or the user's choice is lost. Zero means "not set": the receiving command + // then applies its own flag default. + if o.MaxClients > 0 { + proxyCommand += " --max-clients=" + strconv.Itoa(o.MaxClients) + } + + if o.ServerTimeout > 0 { + proxyCommand += " --server-timeout=" + o.ServerTimeout.String() + } + if o.ServerMetadata != "" { proxyCommand += " --metadata=" + o.ServerMetadata } diff --git a/experimental/ssh/internal/client/client_test.go b/experimental/ssh/internal/client/client_test.go index 09fd1067d4f..5435bf96130 100644 --- a/experimental/ssh/internal/client/client_test.go +++ b/experimental/ssh/internal/client/client_test.go @@ -123,6 +123,81 @@ func TestValidate(t *testing.T) { }, } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + opts := tt.opts + // Validate range-checks MaxClients and ServerTimeout, which the commands always + // supply from a flag default. Set them for every case here so the cases above, + // which are about unrelated fields, aren't rejected by those checks; + // TestValidateServerLifecycle covers them directly. + opts.MaxClients = 10 + opts.ServerTimeout = 24 * time.Hour + err := opts.Validate() + if tt.wantErr == "" { + assert.NoError(t, err) + } else { + assert.EqualError(t, err, tt.wantErr) + } + }) + } +} + +func TestValidateServerLifecycle(t *testing.T) { + tests := []struct { + name string + opts client.ClientOptions + wantErr string + }{ + { + name: "defaults", + opts: client.ClientOptions{ClusterID: "abc-123", MaxClients: 10, ShutdownDelay: 10 * time.Minute, ServerTimeout: 24 * time.Hour}, + }, + { + name: "single client slot", + opts: client.ClientOptions{ClusterID: "abc-123", MaxClients: 1, ServerTimeout: time.Hour}, + }, + { + name: "zero max clients", + opts: client.ClientOptions{ClusterID: "abc-123", MaxClients: 0, ServerTimeout: time.Hour}, + wantErr: "--max-clients must be at least 1, got 0", + }, + { + name: "negative max clients", + opts: client.ClientOptions{ClusterID: "abc-123", MaxClients: -1, ServerTimeout: time.Hour}, + wantErr: "--max-clients must be at least 1, got -1", + }, + { + name: "zero server timeout", + opts: client.ClientOptions{ClusterID: "abc-123", MaxClients: 10}, + wantErr: "--server-timeout must be at least 1s, got 0s", + }, + { + name: "negative server timeout", + opts: client.ClientOptions{ClusterID: "abc-123", MaxClients: 10, ServerTimeout: -time.Minute}, + wantErr: "--server-timeout must be at least 1s, got -1m0s", + }, + { + // Anything under a second truncates to timeout_seconds: 0, which the Jobs API reads + // as "no timeout" - the opposite of the short lifetime this asks for. + name: "sub-second server timeout", + opts: client.ClientOptions{ClusterID: "abc-123", MaxClients: 10, ServerTimeout: 999 * time.Millisecond}, + wantErr: "--server-timeout must be at least 1s, got 999ms", + }, + { + name: "one second server timeout", + opts: client.ClientOptions{ClusterID: "abc-123", MaxClients: 10, ServerTimeout: time.Second}, + }, + { + name: "shutdown delay longer than server timeout", + opts: client.ClientOptions{ClusterID: "abc-123", MaxClients: 10, ShutdownDelay: 48 * time.Hour, ServerTimeout: 24 * time.Hour}, + wantErr: "--shutdown-delay (48h0m0s) cannot be longer than --server-timeout (24h0m0s)", + }, + { + name: "shutdown delay equal to server timeout", + opts: client.ClientOptions{ClusterID: "abc-123", MaxClients: 10, ShutdownDelay: time.Hour, ServerTimeout: time.Hour}, + }, + } + for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { err := tt.opts.Validate() @@ -247,6 +322,18 @@ func TestToProxyCommand(t *testing.T) { opts: client.ClientOptions{ConnectionName: "my-conn", UsagePolicyID: "pol-1", ShutdownDelay: 2 * time.Minute}, want: quoted + " ssh connect --proxy --name=my-conn --shutdown-delay=2m0s --usage-policy-id=pol-1", }, + { + // Both are fixed at server submission, so a host configured by `ssh setup` can only + // carry them through the ProxyCommand. + name: "with server lifecycle flags", + opts: client.ClientOptions{ClusterID: "abc-123", ShutdownDelay: 5 * time.Minute, MaxClients: 25, ServerTimeout: 48 * time.Hour}, + want: quoted + " ssh connect --proxy --cluster=abc-123 --auto-start-cluster=false --shutdown-delay=5m0s --max-clients=25 --server-timeout=48h0m0s", + }, + { + name: "serverless with server lifecycle flags", + opts: client.ClientOptions{ConnectionName: "my-conn", ShutdownDelay: 2 * time.Minute, MaxClients: 25, ServerTimeout: 48 * time.Hour}, + want: quoted + " ssh connect --proxy --name=my-conn --shutdown-delay=2m0s --max-clients=25 --server-timeout=48h0m0s", + }, { name: "with metadata", opts: client.ClientOptions{ClusterID: "abc-123", ServerMetadata: "user,2222,abc-123"}, diff --git a/experimental/ssh/internal/client/submit_internal_test.go b/experimental/ssh/internal/client/submit_internal_test.go index 5b1af006b1d..fa1bd5b2dd1 100644 --- a/experimental/ssh/internal/client/submit_internal_test.go +++ b/experimental/ssh/internal/client/submit_internal_test.go @@ -73,4 +73,47 @@ func TestBuildSSHServerSubmitRun(t *testing.T) { assert.Empty(t, got.Tasks[0].EnvironmentKey) assert.Empty(t, got.Environments) }) + + t.Run("server lifecycle", func(t *testing.T) { + opts := ClientOptions{ + ClusterID: "abc-123", + MaxClients: 25, + ShutdownDelay: 15 * time.Minute, + ServerTimeout: 48 * time.Hour, + } + got := buildSSHServerSubmitRun("v1", "scope", notebookPath, "", opts) + + // This is the only place these two take effect: the server reads maxClients from the + // widget at startup, and the run's timeout caps the tunnel's lifetime. + assert.Equal(t, "25", got.Tasks[0].NotebookTask.BaseParameters["maxClients"]) + assert.Equal(t, "15m0s", got.Tasks[0].NotebookTask.BaseParameters["shutdownDelay"]) + assert.Equal(t, 48*60*60, got.TimeoutSeconds) + assert.Equal(t, 48*60*60, got.Tasks[0].TimeoutSeconds) + }) +} + +// Validate's lower bound on --server-timeout exists only to keep timeout_seconds off 0, which the +// Jobs API reads as "no timeout". Tie the two together: no value Validate accepts may submit an +// unbounded run, whichever side is changed later. +func TestServerTimeoutNeverSubmitsUnboundedRun(t *testing.T) { + const notebookPath = "/Workspace/Users/me/.databricks/ssh-tunnel/v1/conn/ssh-server-bootstrap" + + for _, d := range []time.Duration{ + -time.Minute, + 0, + time.Nanosecond, + 500 * time.Millisecond, + 999 * time.Millisecond, + time.Second, + 10 * time.Minute, + 24 * time.Hour, + } { + opts := ClientOptions{ClusterID: "abc-123", MaxClients: 10, ServerTimeout: d} + if err := opts.Validate(); err != nil { + continue + } + got := buildSSHServerSubmitRun("v1", "scope", notebookPath, "", opts) + assert.NotZero(t, got.TimeoutSeconds, "--server-timeout=%s passed Validate but submits timeout_seconds: 0 (no timeout)", d) + assert.NotZero(t, got.Tasks[0].TimeoutSeconds, "--server-timeout=%s passed Validate but submits a task with timeout_seconds: 0", d) + } } diff --git a/experimental/ssh/internal/setup/setup.go b/experimental/ssh/internal/setup/setup.go index 286ab7be7a7..aa77e87555d 100644 --- a/experimental/ssh/internal/setup/setup.go +++ b/experimental/ssh/internal/setup/setup.go @@ -23,6 +23,13 @@ type SetupOptions struct { AutoStartCluster bool // Delay before shutting down the SSH tunnel, will be added as a --shutdown-delay flag to the ProxyCommand ShutdownDelay time.Duration + // Maximum number of concurrent SSH clients the server accepts, will be added as a --max-clients + // flag to the ProxyCommand. Fixed when the server job is submitted, so the ProxyCommand is the + // only place it can be set for a host configured through setup. + MaxClients int + // Maximum lifetime of the SSH server, will be added as a --server-timeout flag to the ProxyCommand. + // Also fixed at submission time. + ServerTimeout time.Duration // Optional path to the local ssh config. Defaults to ~/.ssh/config SSHConfigPath string // Optional path to the local directory to store SSH keys. Defaults to ~/.databricks/ssh-tunnel-keys @@ -66,6 +73,18 @@ func defaultClusterSelectionPrompt(ctx context.Context, client *databricks.Works } func Setup(ctx context.Context, client *databricks.WorkspaceClient, opts SetupOptions) error { + // Reject invalid server-lifecycle flag values before the cluster picker and + // cluster-access check: these values don't depend on cluster details. + if opts.MaxClients < 1 { + return fmt.Errorf("--max-clients must be at least 1, got %d", opts.MaxClients) + } + if opts.ServerTimeout < time.Second { + return fmt.Errorf("--server-timeout must be at least 1s, got %s", opts.ServerTimeout) + } + if opts.ShutdownDelay > opts.ServerTimeout { + return fmt.Errorf("--shutdown-delay (%s) cannot be longer than --server-timeout (%s)", opts.ShutdownDelay, opts.ServerTimeout) + } + if opts.ClusterID == "" { id, err := clusterSelectionPrompt(ctx, client) if err != nil { @@ -90,8 +109,15 @@ func Setup(ctx context.Context, client *databricks.WorkspaceClient, opts SetupOp ClusterID: opts.ClusterID, AutoStartCluster: opts.AutoStartCluster, ShutdownDelay: opts.ShutdownDelay, + MaxClients: opts.MaxClients, + ServerTimeout: opts.ServerTimeout, Profile: opts.Profile, } + // The ProxyCommand is persisted in the SSH config, so reject values that would produce a + // tunnel that can never work (e.g. --max-clients=0) here rather than at first `ssh `. + if err := clientOpts.Validate(); err != nil { + return err + } proxyCommand, err := clientOpts.ToProxyCommand() if err != nil { return fmt.Errorf("failed to generate ProxyCommand: %w", err) diff --git a/experimental/ssh/internal/setup/setup_test.go b/experimental/ssh/internal/setup/setup_test.go index 9b3267c6b62..7ce1097dcdc 100644 --- a/experimental/ssh/internal/setup/setup_test.go +++ b/experimental/ssh/internal/setup/setup_test.go @@ -183,6 +183,8 @@ func TestSetup_SuccessfulWithNewConfigFile(t *testing.T) { SSHConfigPath: configPath, SSHKeysDir: tmpDir, ShutdownDelay: 30 * time.Second, + MaxClients: 10, + ServerTimeout: 24 * time.Hour, Profile: "test-profile", } @@ -234,6 +236,8 @@ func TestSetup_AutoApproveRecreatesExistingHost(t *testing.T) { SSHConfigPath: configPath, SSHKeysDir: tmpDir, ShutdownDelay: 30 * time.Second, + MaxClients: 10, + ServerTimeout: 24 * time.Hour, AutoApprove: true, } @@ -248,6 +252,86 @@ func TestSetup_AutoApproveRecreatesExistingHost(t *testing.T) { assert.Contains(t, s, "--cluster=cluster-123") } +func TestSetup_SerializesServerLifecycleFlags(t *testing.T) { + ctx := cmdio.MockDiscard(t.Context()) + tmpDir := t.TempDir() + t.Setenv("HOME", tmpDir) + t.Setenv("USERPROFILE", tmpDir) + + m := mocks.NewMockWorkspaceClient(t) + m.GetMockClustersAPI().EXPECT().Get(ctx, compute.GetClusterRequest{ClusterId: "cluster-123"}).Return(&compute.ClusterDetails{ + DataSecurityMode: compute.DataSecurityModeSingleUser, + SingleUserName: "me@example.com", + }, nil) + + opts := SetupOptions{ + HostName: "test-host", + ClusterID: "cluster-123", + SSHConfigPath: filepath.Join(tmpDir, "ssh_config"), + SSHKeysDir: tmpDir, + ShutdownDelay: 30 * time.Second, + MaxClients: 25, + ServerTimeout: 48 * time.Hour, + } + + require.NoError(t, Setup(ctx, m.WorkspaceClient, opts)) + + // The ProxyCommand is the invocation that submits the server job, so both values have to + // reach the persisted host config or the user's choice is silently dropped. + hostContent, err := os.ReadFile(filepath.Join(tmpDir, ".databricks", "ssh-tunnel-configs", "test-host")) + require.NoError(t, err) + assert.Contains(t, string(hostContent), "--max-clients=25") + assert.Contains(t, string(hostContent), "--server-timeout=48h0m0s") +} + +func TestSetup_RejectsUnusableServerLifecycleFlags(t *testing.T) { + tests := []struct { + name string + opts SetupOptions + wantErr string + }{ + { + name: "zero max clients", + opts: SetupOptions{ServerTimeout: 24 * time.Hour}, + wantErr: "--max-clients must be at least 1, got 0", + }, + { + name: "zero server timeout", + opts: SetupOptions{MaxClients: 10}, + wantErr: "--server-timeout must be at least 1s, got 0s", + }, + { + name: "shutdown delay longer than server timeout", + opts: SetupOptions{MaxClients: 10, ShutdownDelay: 48 * time.Hour, ServerTimeout: 24 * time.Hour}, + wantErr: "--shutdown-delay (48h0m0s) cannot be longer than --server-timeout (24h0m0s)", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ctx := cmdio.MockDiscard(t.Context()) + tmpDir := t.TempDir() + t.Setenv("HOME", tmpDir) + t.Setenv("USERPROFILE", tmpDir) + + // Validation fires before any cluster API calls, so no mock expectations needed. + m := mocks.NewMockWorkspaceClient(t) + + opts := tt.opts + opts.HostName = "test-host" + opts.ClusterID = "cluster-123" + opts.SSHConfigPath = filepath.Join(tmpDir, "ssh_config") + opts.SSHKeysDir = tmpDir + + assert.EqualError(t, Setup(ctx, m.WorkspaceClient, opts), tt.wantErr) + + // Nothing is written when the values are rejected. + _, err := os.Stat(filepath.Join(tmpDir, ".databricks", "ssh-tunnel-configs", "test-host")) + assert.ErrorIs(t, err, os.ErrNotExist) + }) + } +} + func TestSetup_PromptsForClusterWhenNotProvided(t *testing.T) { ctx := cmdio.MockDiscard(t.Context()) tmpDir := t.TempDir() @@ -278,6 +362,8 @@ func TestSetup_PromptsForClusterWhenNotProvided(t *testing.T) { SSHConfigPath: configPath, SSHKeysDir: tmpDir, ShutdownDelay: 30 * time.Second, + MaxClients: 10, + ServerTimeout: 24 * time.Hour, } err := Setup(ctx, m.WorkspaceClient, opts) @@ -320,6 +406,8 @@ func TestSetup_SuccessfulWithExistingConfigFile(t *testing.T) { SSHConfigPath: configPath, SSHKeysDir: tmpDir, ShutdownDelay: 60 * time.Second, + MaxClients: 10, + ServerTimeout: 24 * time.Hour, } err = Setup(ctx, m.WorkspaceClient, opts)