mirror of
https://github.com/henrygd/beszel.git
synced 2026-10-02 06:17:47 +02:00
feat(hub): add DISABLE_SSH option to skip SSH fallback (#2469)
Co-authored-by: hank <hank@henrygd.me>
This commit is contained in:
@@ -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 {
|
||||||
|
|||||||
@@ -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")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user