feat(hub): add DISABLE_SSH option to skip SSH fallback (#2469)

Co-authored-by: hank <hank@henrygd.me>
This commit is contained in:
Sven van Ginkel
2026-10-01 17:20:00 +02:00
committed by GitHub
parent 0b101cc6be
commit 8a173e2643
2 changed files with 76 additions and 0 deletions

View File

@@ -66,6 +66,15 @@ type System struct {
lastSavedMonitorProbe map[string]int64 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. // GetStatus returns the current monitoring status.
func (sys *System) GetStatus() string { func (sys *System) GetStatus() string {
sys.statusMu.RLock() sys.statusMu.RLock()
@@ -713,7 +722,11 @@ func shouldCloseWebSocket(err error) bool {
// getSSHTransport returns the system's SSH transport, creating it on first use. // 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 // The transport owns the only SSH connection to the agent; it is shared by the
// updater and on-demand requests and connects lazily. // 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) { func (sys *System) getSSHTransport() (*transport.SSHTransport, error) {
if sshFallbackDisabled() {
return nil, errSSHDisabled
}
sys.sshMu.Lock() sys.sshMu.Lock()
defer sys.sshMu.Unlock() defer sys.sshMu.Unlock()
if sys.sshTransport != nil { if sys.sshTransport != nil {

View File

@@ -4,9 +4,14 @@ package systems
import ( import (
"context" "context"
"net"
"testing" "testing"
"time"
"github.com/henrygd/beszel/internal/common"
"github.com/henrygd/beszel/internal/entities/system" "github.com/henrygd/beszel/internal/entities/system"
"github.com/stretchr/testify/require"
"golang.org/x/crypto/ssh"
) )
func TestCombinedData_MigrateDeprecatedFields(t *testing.T) { 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) 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")
}
})
}
}