diff --git a/internal/hub/systems/system.go b/internal/hub/systems/system.go index f5888f75b..0a292668d 100644 --- a/internal/hub/systems/system.go +++ b/internal/hub/systems/system.go @@ -66,6 +66,15 @@ type System struct { lastSavedMonitorProbe map[string]int64 } +// errSSHDisabled is returned instead of dialing SSH when DISABLE_SSH is set on the hub. +var errSSHDisabled = errors.New("no WebSocket connection and SSH is disabled") + +// sshFallbackDisabled reports whether DISABLE_SSH is set on the hub. +func sshFallbackDisabled() bool { + disableSSH, _ := utils.GetEnv("DISABLE_SSH") + return disableSSH == "true" +} + // GetStatus returns the current monitoring status. func (sys *System) GetStatus() string { sys.statusMu.RLock() @@ -713,7 +722,11 @@ func shouldCloseWebSocket(err error) bool { // getSSHTransport returns the system's SSH transport, creating it on first use. // The transport owns the only SSH connection to the agent; it is shared by the // updater and on-demand requests and connects lazily. +// It returns errSSHDisabled if DISABLE_SSH is set on the hub. func (sys *System) getSSHTransport() (*transport.SSHTransport, error) { + if sshFallbackDisabled() { + return nil, errSSHDisabled + } sys.sshMu.Lock() defer sys.sshMu.Unlock() if sys.sshTransport != nil { diff --git a/internal/hub/systems/system_test.go b/internal/hub/systems/system_test.go index ab720c7a5..e7eac097e 100644 --- a/internal/hub/systems/system_test.go +++ b/internal/hub/systems/system_test.go @@ -4,9 +4,14 @@ package systems import ( "context" + "net" "testing" + "time" + "github.com/henrygd/beszel/internal/common" "github.com/henrygd/beszel/internal/entities/system" + "github.com/stretchr/testify/require" + "golang.org/x/crypto/ssh" ) func TestCombinedData_MigrateDeprecatedFields(t *testing.T) { @@ -172,3 +177,61 @@ func TestSetDownAfterContextCancelled(t *testing.T) { t.Fatalf("status should be untouched, got %q", sys.Status) } } + +func TestSSHDisabledSkipsSSHFallback(t *testing.T) { + t.Setenv("DISABLE_SSH", "true") + // manager is nil on purpose: any SSH attempt would panic + sys := &System{} + + var result string + err := sys.request(context.Background(), common.GetContainerInfo, nil, &result) + require.ErrorIs(t, err, errSSHDisabled) + + _, err = sys.fetchDataFromAgent(common.DataRequestOptions{}) + require.ErrorIs(t, err, errSSHDisabled) +} + +func TestSSHFallbackDialsAgentUnlessDisabled(t *testing.T) { + for _, tc := range []struct { + name string + hubEnv string + wantDial bool + }{ + {name: "enabled", wantDial: true}, + {name: "disabled on hub", hubEnv: "DISABLE_SSH"}, + {name: "disabled on hub with prefix", hubEnv: "BESZEL_HUB_DISABLE_SSH"}, + } { + t.Run(tc.name, func(t *testing.T) { + if tc.hubEnv != "" { + t.Setenv(tc.hubEnv, "true") + } + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + defer listener.Close() + accepted := make(chan struct{}, 1) + go func() { + if conn, err := listener.Accept(); err == nil { + accepted <- struct{}{} + conn.Close() + } + }() + + host, port, _ := net.SplitHostPort(listener.Addr().String()) + sm := &SystemManager{sshConfig: &ssh.ClientConfig{HostKeyCallback: ssh.InsecureIgnoreHostKey(), Timeout: time.Second}} + sys := &System{Host: host, Port: port, Status: down, manager: sm, ctx: context.Background()} + + _, err = sys.fetchDataFromAgent(common.DataRequestOptions{}) + require.Error(t, err) // listener isn't a real agent + if !tc.wantDial { + require.ErrorIs(t, err, errSSHDisabled) + } + + select { + case <-accepted: + require.True(t, tc.wantDial, "hub must not dial SSH when it is disabled") + case <-time.After(200 * time.Millisecond): + require.False(t, tc.wantDial, "hub must dial SSH when it is enabled") + } + }) + } +}