Files
beszel-ipv6/agent/server_test.go
spatiumstas 4ffd83677d fix(agent): improve WebSocket and SSH fallback handling (#2441)
Make WebSocket reconnect and SSH fallback transitions reliable across asynchronous disconnects, stale callbacks, and overlapping connections. Keep the SSH listener available while disconnected and allow a verified WebSocket connection to take precedence when it recovers.

Co-authored-by: henrygd <hank@henrygd.me>
2026-10-01 13:37:57 -04:00

975 lines
28 KiB
Go

//go:build testing
package agent
import (
"context"
"crypto/ed25519"
"encoding/json"
"fmt"
"io"
"net"
"net/http"
"net/http/httptest"
"net/url"
"os"
"path/filepath"
"strings"
"sync"
"testing"
"time"
"github.com/henrygd/beszel/internal/entities/container"
"github.com/henrygd/beszel/internal/entities/system"
"github.com/blang/semver"
"github.com/fxamacker/cbor/v2"
"github.com/gliderlabs/ssh"
"github.com/lxzan/gws"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
gossh "golang.org/x/crypto/ssh"
)
func TestStartServer(t *testing.T) {
// Generate a test key pair
pubKey, privKey, err := ed25519.GenerateKey(nil)
require.NoError(t, err)
signer, err := gossh.NewSignerFromKey(privKey)
require.NoError(t, err)
sshPubKey, err := gossh.NewPublicKey(pubKey)
require.NoError(t, err)
// Generate a different key pair for bad key test
badPubKey, badPrivKey, err := ed25519.GenerateKey(nil)
require.NoError(t, err)
badSigner, err := gossh.NewSignerFromKey(badPrivKey)
require.NoError(t, err)
sshBadPubKey, err := gossh.NewPublicKey(badPubKey)
require.NoError(t, err)
socketFile := filepath.Join(t.TempDir(), "beszel-test.sock")
tests := []struct {
name string
config ServerOptions
wantErr bool
errContains string
setup func() error
cleanup func() error
}{
{
name: "tcp port only",
config: ServerOptions{
Network: "tcp",
Addr: ":45987",
Keys: []gossh.PublicKey{sshPubKey},
},
},
{
name: "tcp with ipv4",
config: ServerOptions{
Network: "tcp4",
Addr: "127.0.0.1:45988",
Keys: []gossh.PublicKey{sshPubKey},
},
},
{
name: "tcp with ipv6",
config: ServerOptions{
Network: "tcp6",
Addr: "[::1]:45989",
Keys: []gossh.PublicKey{sshPubKey},
},
},
{
name: "unix socket",
config: ServerOptions{
Network: "unix",
Addr: socketFile,
Keys: []gossh.PublicKey{sshPubKey},
},
setup: func() error {
// Create a socket file that should be removed
f, err := os.Create(socketFile)
if err != nil {
return err
}
return f.Close()
},
cleanup: func() error {
return os.Remove(socketFile)
},
},
{
name: "bad key should fail",
config: ServerOptions{
Network: "tcp",
Addr: ":45987",
Keys: []gossh.PublicKey{sshBadPubKey},
},
wantErr: true,
errContains: "ssh: handshake failed",
},
{
name: "good key still good",
config: ServerOptions{
Network: "tcp",
Addr: ":45987",
Keys: []gossh.PublicKey{sshPubKey},
},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if tt.setup != nil {
err := tt.setup()
require.NoError(t, err)
}
if tt.cleanup != nil {
defer tt.cleanup()
}
agent, err := NewAgent("")
require.NoError(t, err)
// Start server in a goroutine since it blocks
errChan := make(chan error, 1)
go func() {
errChan <- agent.StartServer(tt.config)
}()
// Add a short delay to allow the server to start
time.Sleep(100 * time.Millisecond)
// Try to connect to verify server is running
var client *gossh.Client
// Choose the appropriate signer based on the test case
testSigner := signer
if tt.name == "bad key should fail" {
testSigner = badSigner
}
sshClientConfig := &gossh.ClientConfig{
User: "a",
Auth: []gossh.AuthMethod{
gossh.PublicKeys(testSigner),
},
HostKeyCallback: gossh.InsecureIgnoreHostKey(),
Timeout: 4 * time.Second,
}
switch tt.config.Network {
case "unix":
client, err = gossh.Dial("unix", tt.config.Addr, sshClientConfig)
default:
if !strings.Contains(tt.config.Addr, ":") {
tt.config.Addr = ":" + tt.config.Addr
}
client, err = gossh.Dial("tcp", tt.config.Addr, sshClientConfig)
}
if tt.wantErr {
assert.Error(t, err)
if tt.errContains != "" {
assert.Contains(t, err.Error(), tt.errContains)
}
return
}
require.NoError(t, err)
require.NotNil(t, client)
client.Close()
})
}
}
func TestStartServerDisableSSH(t *testing.T) {
t.Setenv("BESZEL_AGENT_DISABLE_SSH", "true")
agent, err := NewAgent("")
require.NoError(t, err)
opts := ServerOptions{
Network: "tcp",
Addr: ":45990",
}
err = agent.StartServer(opts)
assert.Error(t, err)
assert.Contains(t, err.Error(), "SSH disabled")
}
func TestStopServerDoesNotBlockWhenEventQueueFull(t *testing.T) {
agent := createTestAgent(t)
agent.server = &ssh.Server{}
agent.connectionManager.eventChan = make(chan ConnectionEvent, 1)
agent.connectionManager.eventChan <- WebSocketConnect
done := make(chan error, 1)
go func() {
done <- agent.StopServer()
}()
select {
case err := <-done:
require.NoError(t, err)
case <-time.After(time.Second):
t.Fatal("StopServer blocked on the connection event queue")
}
assert.Nil(t, agent.server)
assert.Equal(t, WebSocketConnect, <-agent.connectionManager.eventChan)
}
func TestSSHConnectionFallbackLifecycle(t *testing.T) {
t.Setenv("BESZEL_AGENT_DISABLE_SSH", "false")
agent := createTestAgent(t)
cm := agent.connectionManager
cm.eventChan = make(chan ConnectionEvent, 4)
_, privateKey, err := ed25519.GenerateKey(nil)
require.NoError(t, err)
signer, err := gossh.NewSignerFromKey(privateKey)
require.NoError(t, err)
cm.serverOptions = ServerOptions{
Network: "tcp",
Addr: "127.0.0.1:0",
Keys: []gossh.PublicKey{signer.PublicKey()},
}
// A WebSocket that closed after its upgrade returned nil must start SSH.
cm.handleEvent(WebSocketDisconnect)
agent.serverMu.Lock()
require.NotNil(t, agent.serverListener)
addr := agent.serverListener.Addr().String()
agent.serverMu.Unlock()
defer func() { _ = agent.StopServer() }()
clientConfig := &gossh.ClientConfig{
User: "hub",
Auth: []gossh.AuthMethod{gossh.PublicKeys(signer)},
HostKeyCallback: gossh.InsecureIgnoreHostKey(),
Timeout: 4 * time.Second,
}
client, err := gossh.Dial("tcp", addr, clientConfig)
require.NoError(t, err)
defer client.Close()
// A connection is counted when it starts its first session.
startSession := func(c *gossh.Client) *gossh.Session {
session, err := c.NewSession()
require.NoError(t, err)
require.NoError(t, session.Shell())
return session
}
session := startSession(client)
select {
case <-cm.sshChanged:
cm.handleSSHChange()
case <-time.After(5 * time.Second):
t.Fatal("SSH connection did not notify the manager")
}
require.Equal(t, SSHConnected, cm.getState())
wsAttempt := make(chan struct{}, 1)
releaseWS := make(chan struct{})
var releaseOnce sync.Once
release := func() { releaseOnce.Do(func() { close(releaseWS) }) }
hub := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
wsAttempt <- struct{}{}
<-releaseWS
w.WriteHeader(http.StatusServiceUnavailable)
}))
defer hub.Close()
defer release()
t.Setenv("BESZEL_AGENT_HUB_URL", hub.URL)
t.Setenv("BESZEL_AGENT_TOKEN", "test-token")
cm.wsClient, err = newWebSocketClient(agent)
require.NoError(t, err)
// A normal short-lived session must not be mistaken for a lost connection,
// and further sessions must not count the same connection again.
_ = session.Close()
_ = startSession(client).Close()
select {
case <-cm.sshChanged:
t.Fatal("session close unexpectedly changed SSH connection state")
case <-time.After(100 * time.Millisecond):
}
cm.mu.Lock()
assert.Equal(t, 1, cm.sshConnections, "sessions should not be counted as connections")
cm.mu.Unlock()
secondClient, err := gossh.Dial("tcp", addr, clientConfig)
require.NoError(t, err)
defer secondClient.Close()
defer startSession(secondClient).Close()
require.Eventually(t, func() bool {
cm.mu.Lock()
defer cm.mu.Unlock()
return cm.sshConnections == 2
}, 5*time.Second, 10*time.Millisecond, "second SSH connection was not counted")
require.NoError(t, client.Close())
select {
case <-cm.sshChanged:
t.Fatal("closing one of two SSH connections changed SSH connection state")
case <-time.After(100 * time.Millisecond):
}
require.Equal(t, SSHConnected, cm.getState())
require.NoError(t, secondClient.Close())
select {
case <-cm.sshChanged:
cm.handleSSHChange()
case <-time.After(5 * time.Second):
t.Fatal("SSH TCP close did not notify the manager")
}
require.Equal(t, Disconnected, cm.getState())
require.NotNil(t, cm.wsTicker)
select {
case <-wsAttempt:
case <-time.After(5 * time.Second):
t.Fatal("agent did not retry WebSocket after SSH disconnected")
}
// The hub may redial straight away, so the listener must stay open while
// the WebSocket attempt is pending and after it fails.
requireSameListener := func(msg string) {
agent.serverMu.Lock()
defer agent.serverMu.Unlock()
require.NotNil(t, agent.serverListener, msg)
assert.Equal(t, addr, agent.serverListener.Addr().String(), msg)
}
requireSameListener("SSH listener should stay open during the WebSocket attempt")
thirdClient, err := gossh.Dial("tcp", addr, clientConfig)
require.NoError(t, err, "SSH should accept a redial during the WebSocket attempt")
require.NoError(t, thirdClient.Close())
release()
require.Eventually(t, func() bool {
return !cm.isConnectingNow()
}, 5*time.Second, 10*time.Millisecond, "reconnect attempt did not finish")
requireSameListener("SSH listener should stay open after the WebSocket attempt fails")
cm.stopWsTicker()
}
// offeredKeySigner offers an authorized public key without proving possession
// of its private key: Sign blocks until released, then signs with another key.
type offeredKeySigner struct {
gossh.Signer
publicKey gossh.PublicKey
signing chan struct{}
release chan struct{}
}
func (s *offeredKeySigner) PublicKey() gossh.PublicKey { return s.publicKey }
func (s *offeredKeySigner) Sign(rand io.Reader, data []byte) (*gossh.Signature, error) {
close(s.signing)
<-s.release
return s.Signer.Sign(rand, data)
}
// The public key handler runs when a key is offered, before the client signs
// anything, so it must not be what marks an SSH connection as established.
func TestSSHPublicKeyOfferIsNotAConnection(t *testing.T) {
t.Setenv("BESZEL_AGENT_DISABLE_SSH", "false")
agent := createTestAgent(t)
cm := agent.connectionManager
cm.eventChan = make(chan ConnectionEvent, 4)
newSigner := func() gossh.Signer {
_, privateKey, err := ed25519.GenerateKey(nil)
require.NoError(t, err)
signer, err := gossh.NewSignerFromKey(privateKey)
require.NoError(t, err)
return signer
}
hubKey := newSigner().PublicKey()
cm.serverOptions = ServerOptions{
Network: "tcp",
Addr: "127.0.0.1:0",
Keys: []gossh.PublicKey{hubKey},
}
cm.handleEvent(WebSocketDisconnect)
agent.serverMu.Lock()
require.NotNil(t, agent.serverListener)
addr := agent.serverListener.Addr().String()
agent.serverMu.Unlock()
defer func() { _ = agent.StopServer() }()
signer := &offeredKeySigner{
Signer: newSigner(),
publicKey: hubKey,
signing: make(chan struct{}),
release: make(chan struct{}),
}
dialErr := make(chan error, 1)
go func() {
client, err := gossh.Dial("tcp", addr, &gossh.ClientConfig{
User: "hub",
Auth: []gossh.AuthMethod{gossh.PublicKeys(signer)},
HostKeyCallback: gossh.InsecureIgnoreHostKey(),
Timeout: 4 * time.Second,
})
if client != nil {
client.Close()
}
dialErr <- err
}()
// The server has accepted the offered key and is waiting for a signature.
select {
case <-signer.signing:
case <-time.After(5 * time.Second):
t.Fatal("server did not accept the offered public key")
}
select {
case <-cm.sshChanged:
t.Fatal("offering a public key changed SSH connection state")
case <-time.After(100 * time.Millisecond):
}
assert.False(t, cm.hasSSHConnection())
assert.Equal(t, Disconnected, cm.getState())
close(signer.release)
require.Error(t, <-dialErr, "a signature from another key must be rejected")
assert.False(t, cm.hasSSHConnection())
}
// startSSHFallbackServer starts the fallback SSH server for a disconnected
// agent and returns its address and a client config that can authenticate.
func startSSHFallbackServer(t *testing.T) (*Agent, string, *gossh.ClientConfig) {
t.Helper()
t.Setenv("BESZEL_AGENT_DISABLE_SSH", "false")
agent := createTestAgent(t)
cm := agent.connectionManager
cm.eventChan = make(chan ConnectionEvent, 4)
_, privateKey, err := ed25519.GenerateKey(nil)
require.NoError(t, err)
signer, err := gossh.NewSignerFromKey(privateKey)
require.NoError(t, err)
cm.serverOptions = ServerOptions{
Network: "tcp",
Addr: "127.0.0.1:0",
Keys: []gossh.PublicKey{signer.PublicKey()},
}
cm.startSSHServer()
agent.serverMu.Lock()
require.NotNil(t, agent.serverListener)
addr := agent.serverListener.Addr().String()
agent.serverMu.Unlock()
t.Cleanup(func() { _ = agent.StopServer() })
return agent, addr, &gossh.ClientConfig{
User: "hub",
Auth: []gossh.AuthMethod{gossh.PublicKeys(signer)},
HostKeyCallback: gossh.InsecureIgnoreHostKey(),
Timeout: 4 * time.Second,
}
}
// connectSSHSession dials the agent and starts a session, which is what marks
// the connection as established.
func connectSSHSession(t *testing.T, addr string, config *gossh.ClientConfig) *gossh.Client {
t.Helper()
client, err := gossh.Dial("tcp", addr, config)
require.NoError(t, err)
t.Cleanup(func() { _ = client.Close() })
session, err := client.NewSession()
require.NoError(t, err)
require.NoError(t, session.Shell())
return client
}
// handleNextSSHChange applies the next SSH connection notification, as the
// connection manager's event loop would.
func handleNextSSHChange(t *testing.T, cm *ConnectionManager) {
t.Helper()
select {
case <-cm.sshChanged:
cm.handleSSHChange()
case <-time.After(5 * time.Second):
t.Fatal("SSH connection change did not notify the manager")
}
}
// An agent without a WebSocket client only has SSH, so losing the hub's SSH
// connection must leave the listener in place for it to reconnect.
func TestSSHDisconnectKeepsListenerWithoutWebSocket(t *testing.T) {
agent, addr, clientConfig := startSSHFallbackServer(t)
cm := agent.connectionManager
require.Nil(t, cm.wsClient)
defer cm.stopWsTicker()
client := connectSSHSession(t, addr, clientConfig)
handleNextSSHChange(t, cm)
require.Equal(t, SSHConnected, cm.getState())
require.NoError(t, client.Close())
handleNextSSHChange(t, cm)
require.Equal(t, Disconnected, cm.getState())
require.Eventually(t, func() bool {
return !cm.isConnectingNow()
}, 5*time.Second, 10*time.Millisecond, "reconnect attempt did not finish")
agent.serverMu.Lock()
require.NotNil(t, agent.serverListener, "SSH listener should stay open")
assert.Equal(t, addr, agent.serverListener.Addr().String(), "SSH listener should not be restarted")
agent.serverMu.Unlock()
connectSSHSession(t, addr, clientConfig)
handleNextSSHChange(t, cm)
assert.Equal(t, SSHConnected, cm.getState())
}
// A WebSocket attempt that was already in flight can authenticate after SSH has
// connected. WebSocket is preferred, so it takes over and SSH is shut down.
func TestWebSocketTakesOverFromSSH(t *testing.T) {
agent, addr, clientConfig := startSSHFallbackServer(t)
cm := agent.connectionManager
client := connectSSHSession(t, addr, clientConfig)
handleNextSSHChange(t, cm)
require.Equal(t, SSHConnected, cm.getState())
cm.wsClient = &WebSocketClient{
agent: agent,
hubURL: &url.URL{Host: "localhost:8080"},
Conn: &gws.Conn{},
hubVerified: true,
}
cm.handleEvent(WebSocketConnect)
require.Equal(t, WebSocketConnected, cm.getState())
agent.serverMu.Lock()
assert.Nil(t, agent.serverListener, "SSH listener should close once WebSocket takes over")
agent.serverMu.Unlock()
closed := make(chan struct{})
go func() {
_ = client.Wait()
close(closed)
}()
select {
case <-closed:
case <-time.After(5 * time.Second):
t.Fatal("SSH connection was not closed when WebSocket took over")
}
// The SSH connection closing must not disturb the WebSocket state.
handleNextSSHChange(t, cm)
assert.False(t, cm.hasSSHConnection())
assert.Equal(t, WebSocketConnected, cm.getState())
}
/////////////////////////////////////////////////////////////////
//////////////////// ParseKeys Tests ////////////////////////////
/////////////////////////////////////////////////////////////////
// Helper function to generate a temporary file with content
func createTempFile(content string) (string, error) {
tmpFile, err := os.CreateTemp("", "ssh_keys_*.txt")
if err != nil {
return "", fmt.Errorf("failed to create temp file: %w", err)
}
defer tmpFile.Close()
if _, err := tmpFile.WriteString(content); err != nil {
return "", fmt.Errorf("failed to write to temp file: %w", err)
}
return tmpFile.Name(), nil
}
// Test case 1: String with a single SSH key
func TestParseSingleKeyFromString(t *testing.T) {
input := "ssh-ed25519 AAAAC3NzaC1lZDI1NTE5AAAAIKCBM91kukN7hbvFKtbpEeo2JXjCcNxXcdBH7V7ADMBo"
keys, err := ParseKeys(input)
if err != nil {
t.Fatalf("Expected no error, got: %v", err)
}
if len(keys) != 1 {
t.Fatalf("Expected 1 key, got %d keys", len(keys))
}
if keys[0].Type() != "ssh-ed25519" {
t.Fatalf("Expected key type 'ssh-ed25519', got '%s'", keys[0].Type())
}
}
// Test case 2: String with multiple SSH keys
func TestParseMultipleKeysFromString(t *testing.T) {
input := "ssh-ed25519 AAAAC3NzaC1lZDI1NTE5AAAAIKCBM91kukN7hbvFKtbpEeo2JXjCcNxXcdBH7V7ADMBo\nssh-ed25519 AAAAC3NzaC1lZDI1NTE5AAAAIJDMtAOQfxDlCxe+A5lVbUY/DHxK1LAF2Z3AV0FYv36D \n #comment\n ssh-ed25519 AAAAC3NzaC1lZDI1NTE5AAAAIJDMtAOQfxDlCxe+A5lVbUY/DHxK1LAF2Z3AV0FYv36D"
keys, err := ParseKeys(input)
if err != nil {
t.Fatalf("Expected no error, got: %v", err)
}
if len(keys) != 3 {
t.Fatalf("Expected 3 keys, got %d keys", len(keys))
}
if keys[0].Type() != "ssh-ed25519" || keys[1].Type() != "ssh-ed25519" || keys[2].Type() != "ssh-ed25519" {
t.Fatalf("Unexpected key types: %s, %s, %s", keys[0].Type(), keys[1].Type(), keys[2].Type())
}
}
// Test case 3: File with a single SSH key
func TestParseSingleKeyFromFile(t *testing.T) {
content := "ssh-ed25519 AAAAC3NzaC1lZDI1NTE5AAAAIKCBM91kukN7hbvFKtbpEeo2JXjCcNxXcdBH7V7ADMBo"
filePath, err := createTempFile(content)
if err != nil {
t.Fatalf("Failed to create temp file: %v", err)
}
defer os.Remove(filePath) // Clean up the file after the test
// Read the file content
fileContent, err := os.ReadFile(filePath)
if err != nil {
t.Fatalf("Failed to read temp file: %v", err)
}
// Parse the keys
keys, err := ParseKeys(string(fileContent))
if err != nil {
t.Fatalf("Expected no error, got: %v", err)
}
if len(keys) != 1 {
t.Fatalf("Expected 1 key, got %d keys", len(keys))
}
if keys[0].Type() != "ssh-ed25519" {
t.Fatalf("Expected key type 'ssh-ed25519', got '%s'", keys[0].Type())
}
}
// Test case 4: File with multiple SSH keys
func TestParseMultipleKeysFromFile(t *testing.T) {
content := "ssh-ed25519 AAAAC3NzaC1lZDI1NTE5AAAAIKCBM91kukN7hbvFKtbpEeo2JXjCcNxXcdBH7V7ADMBo\nssh-ed25519 AAAAC3NzaC1lZDI1NTE5AAAAIJDMtAOQfxDlCxe+A5lVbUY/DHxK1LAF2Z3AV0FYv36D \n #comment\n ssh-ed25519 AAAAC3NzaC1lZDI1NTE5AAAAIJDMtAOQfxDlCxe+A5lVbUY/DHxK1LAF2Z3AV0FYv36D"
filePath, err := createTempFile(content)
if err != nil {
t.Fatalf("Failed to create temp file: %v", err)
}
// defer os.Remove(filePath) // Clean up the file after the test
// Read the file content
fileContent, err := os.ReadFile(filePath)
if err != nil {
t.Fatalf("Failed to read temp file: %v", err)
}
// Parse the keys
keys, err := ParseKeys(string(fileContent))
if err != nil {
t.Fatalf("Expected no error, got: %v", err)
}
if len(keys) != 3 {
t.Fatalf("Expected 3 keys, got %d keys", len(keys))
}
if keys[0].Type() != "ssh-ed25519" || keys[1].Type() != "ssh-ed25519" || keys[2].Type() != "ssh-ed25519" {
t.Fatalf("Unexpected key types: %s, %s, %s", keys[0].Type(), keys[1].Type(), keys[2].Type())
}
}
// Test case 5: Invalid SSH key input
func TestParseInvalidKey(t *testing.T) {
input := "invalid-key-data"
_, err := ParseKeys(input)
if err == nil {
t.Fatalf("Expected an error for invalid key, got nil")
}
expectedErrMsg := "failed to parse key"
if !strings.Contains(err.Error(), expectedErrMsg) {
t.Fatalf("Expected error message to contain '%s', got: %v", expectedErrMsg, err)
}
}
/////////////////////////////////////////////////////////////////
//////////////////// Hub Version Tests //////////////////////////
/////////////////////////////////////////////////////////////////
func TestExtractHubVersion(t *testing.T) {
tests := []struct {
name string
clientVersion string
expectedVersion string
expectError bool
}{
{
name: "valid beszel client version with underscore",
clientVersion: "SSH-2.0-beszel_0.11.1",
expectedVersion: "0.11.1",
expectError: false,
},
{
name: "valid beszel client version with beta",
clientVersion: "SSH-2.0-beszel_1.0.0-beta",
expectedVersion: "1.0.0-beta",
expectError: false,
},
{
name: "valid beszel client version with rc",
clientVersion: "SSH-2.0-beszel_0.12.0-rc1",
expectedVersion: "0.12.0-rc1",
expectError: false,
},
{
name: "different SSH client",
clientVersion: "SSH-2.0-OpenSSH_8.0",
expectedVersion: "8.0",
expectError: true,
},
{
name: "malformed version string without underscore",
clientVersion: "SSH-2.0-beszel",
expectError: true,
},
{
name: "empty version string",
clientVersion: "",
expectError: true,
},
{
name: "version string with underscore but no version",
clientVersion: "beszel_",
expectedVersion: "",
expectError: true,
},
{
name: "version with patch and build metadata",
clientVersion: "SSH-2.0-beszel_1.2.3+build.123",
expectedVersion: "1.2.3+build.123",
expectError: false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
result, err := extractHubVersion(tt.clientVersion)
if tt.expectError {
assert.Error(t, err)
return
}
require.NoError(t, err)
assert.Equal(t, tt.expectedVersion, result.String())
})
}
}
/////////////////////////////////////////////////////////////////
/////////////// Hub Version Detection Tests ////////////////////
/////////////////////////////////////////////////////////////////
func TestGetHubVersion(t *testing.T) {
agent, err := NewAgent("")
require.NoError(t, err)
// Mock SSH context that implements the ssh.Context interface
mockCtx := &mockSSHContext{
sessionID: "test-session-123",
clientVersion: "SSH-2.0-beszel_0.12.0",
}
// Test first call - should extract version
version := agent.getHubVersion(mockCtx)
assert.Equal(t, "0.12.0", version.String())
// Test that version reflects the current client version (no stale caching)
mockCtx.clientVersion = "SSH-2.0-beszel_0.11.0"
version = agent.getHubVersion(mockCtx)
assert.Equal(t, "0.11.0", version.String())
// Test with invalid version string (non-beszel client)
mockCtx.clientVersion = "SSH-2.0-OpenSSH_8.0"
version = agent.getHubVersion(mockCtx)
assert.Equal(t, "0.0.0", version.String()) // Should be empty version for non-beszel clients
// Test with no client version
mockCtx.clientVersion = ""
version = agent.getHubVersion(mockCtx)
assert.True(t, version.EQ(semver.Version{})) // Should be empty version
}
// mockSSHContext implements ssh.Context for testing
type mockSSHContext struct {
context.Context
sync.Mutex
sessionID string
clientVersion string
}
func (m *mockSSHContext) SessionID() string {
return m.sessionID
}
func (m *mockSSHContext) ClientVersion() string {
return m.clientVersion
}
func (m *mockSSHContext) ServerVersion() string {
return "SSH-2.0-beszel_test"
}
func (m *mockSSHContext) Value(key interface{}) interface{} {
if key == ssh.ContextKeyClientVersion {
return m.clientVersion
}
return nil
}
func (m *mockSSHContext) User() string { return "test-user" }
func (m *mockSSHContext) RemoteAddr() net.Addr { return nil }
func (m *mockSSHContext) LocalAddr() net.Addr { return nil }
func (m *mockSSHContext) Permissions() *ssh.Permissions { return nil }
func (m *mockSSHContext) SetValue(key, value interface{}) {}
/////////////////////////////////////////////////////////////////
/////////////// CBOR vs JSON Encoding Tests ////////////////////
/////////////////////////////////////////////////////////////////
// TestWriteToSessionEncoding tests that writeToSession actually encodes data in the correct format
func TestWriteToSessionEncoding(t *testing.T) {
tests := []struct {
name string
hubVersion string
expectedUsesCbor bool
}{
{
name: "old hub version should use JSON",
hubVersion: "0.11.1",
expectedUsesCbor: false,
},
{
name: "non-beta release should use CBOR",
hubVersion: "0.12.0",
expectedUsesCbor: true,
},
{
name: "even newer hub version should use CBOR",
hubVersion: "0.16.4",
expectedUsesCbor: true,
},
{
name: "beta version below release threshold should use JSON",
hubVersion: "0.12.0-beta0",
expectedUsesCbor: false,
},
// {
// name: "matching beta version should use CBOR",
// hubVersion: "0.12.0-beta2",
// expectedUsesCbor: true,
// },
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
agent, err := NewAgent("")
require.NoError(t, err)
// Parse the test version
version, err := semver.Parse(tt.hubVersion)
require.NoError(t, err)
// Create test data to encode
testData := createTestCombinedData()
var buf strings.Builder
err = agent.writeToSession(&buf, testData, version)
require.NoError(t, err)
encodedData := buf.String()
require.NotEmpty(t, encodedData)
// Verify the encoding format by attempting to decode
if tt.expectedUsesCbor {
var decodedCbor system.CombinedData
err = cbor.Unmarshal([]byte(encodedData), &decodedCbor)
assert.NoError(t, err, "Should be valid CBOR data")
var decodedJson system.CombinedData
err = json.Unmarshal([]byte(encodedData), &decodedJson)
assert.Error(t, err, "Should not be valid JSON data")
assert.Equal(t, testData.Details.Hostname, decodedCbor.Details.Hostname)
assert.Equal(t, testData.Stats.Cpu, decodedCbor.Stats.Cpu)
} else {
// Should be JSON - try to decode as JSON
var decodedJson system.CombinedData
err = json.Unmarshal([]byte(encodedData), &decodedJson)
assert.NoError(t, err, "Should be valid JSON data")
var decodedCbor system.CombinedData
err = cbor.Unmarshal([]byte(encodedData), &decodedCbor)
assert.Error(t, err, "Should not be valid CBOR data")
// Verify the decoded JSON data matches our test data
assert.Equal(t, testData.Details.Hostname, decodedJson.Details.Hostname)
assert.Equal(t, testData.Stats.Cpu, decodedJson.Stats.Cpu)
// Verify it looks like JSON (starts with '{' and contains readable field names)
assert.True(t, strings.HasPrefix(encodedData, "{"), "JSON should start with '{'")
assert.Contains(t, encodedData, `"info"`, "JSON should contain readable field names")
assert.Contains(t, encodedData, `"stats"`, "JSON should contain readable field names")
}
})
}
}
// Helper function to create test data for encoding tests
func createTestCombinedData() *system.CombinedData {
return &system.CombinedData{
Stats: system.Stats{
Cpu: 25.5,
Mem: 8589934592, // 8GB
MemUsed: 4294967296, // 4GB
MemPct: 50.0,
DiskTotal: 1099511627776, // 1TB
DiskUsed: 549755813888, // 512GB
DiskPct: 50.0,
},
Details: &system.Details{
Hostname: "test-host",
},
Info: system.Info{
Uptime: 3600,
AgentVersion: "0.12.0",
},
Containers: []*container.Stats{
{
Name: "test-container",
Cpu: 10.5,
Mem: 1073741824, // 1GB
},
},
}
}
// TestGetHubVersionConcurrent guards against a regression of the
// "concurrent map writes" panic previously caused by a shared, unsynchronized
// hubVersions cache (see https://github.com/henrygd/beszel/issues/2128).
// getHubVersion no longer shares mutable state between sessions, so calling
// it concurrently from many goroutines must be safe under `go test -race`.
func TestGetHubVersionConcurrent(t *testing.T) {
agent, err := NewAgent("")
require.NoError(t, err)
const goroutines = 50
var wg sync.WaitGroup
wg.Add(goroutines)
for i := 0; i < goroutines; i++ {
go func(i int) {
defer wg.Done()
ctx := &mockSSHContext{
sessionID: fmt.Sprintf("session-%d", i),
clientVersion: "SSH-2.0-beszel_0.12.0",
}
version := agent.getHubVersion(ctx)
assert.Equal(t, "0.12.0", version.String())
}(i)
}
wg.Wait()
}