mirror of
https://github.com/henrygd/beszel.git
synced 2026-09-16 21:14:23 +00:00
Merge commit from fork
This commit is contained in:
+44
-1
@@ -2,6 +2,7 @@ package agent
|
||||
|
||||
import (
|
||||
"crypto/tls"
|
||||
"crypto/x509"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
@@ -27,6 +28,18 @@ const (
|
||||
wsDeadline = 70 * time.Second
|
||||
)
|
||||
|
||||
type caCertFileError struct {
|
||||
err error
|
||||
}
|
||||
|
||||
func (e *caCertFileError) Error() string {
|
||||
return e.err.Error()
|
||||
}
|
||||
|
||||
func (e *caCertFileError) Unwrap() error {
|
||||
return e.err
|
||||
}
|
||||
|
||||
// WebSocketClient manages the WebSocket connection between the agent and hub.
|
||||
// It handles authentication, message routing, and connection lifecycle management.
|
||||
type WebSocketClient struct {
|
||||
@@ -40,6 +53,7 @@ type WebSocketClient struct {
|
||||
hubRequest *common.HubRequest[cbor.RawMessage] // Reusable request structure for message parsing
|
||||
lastConnectAttempt time.Time // Timestamp of last connection attempt
|
||||
hubVerified bool // Whether the hub has been cryptographically verified
|
||||
tlsConfig *tls.Config // Optional TLS configuration with custom CA certificates
|
||||
}
|
||||
|
||||
// newWebSocketClient creates a new WebSocket client for the given agent.
|
||||
@@ -61,6 +75,10 @@ func newWebSocketClient(agent *Agent) (client *WebSocketClient, err error) {
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
client.tlsConfig, err = getTLSConfig()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
client.agent = agent
|
||||
client.hubRequest = &common.HubRequest[cbor.RawMessage]{}
|
||||
@@ -110,6 +128,31 @@ func parseTokenFile(contents, path string) (string, error) {
|
||||
return token, nil
|
||||
}
|
||||
|
||||
// getTLSConfig returns a TLS configuration containing the system certificate
|
||||
// pool plus any certificates configured through CA_CERT_FILE. A nil config lets
|
||||
// gws use Go's default TLS configuration and system roots.
|
||||
func getTLSConfig() (*tls.Config, error) {
|
||||
caCertFile, _ := utils.GetEnv("CA_CERT_FILE")
|
||||
if caCertFile == "" {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
caCertPEM, err := os.ReadFile(caCertFile)
|
||||
if err != nil {
|
||||
return nil, &caCertFileError{fmt.Errorf("read CA_CERT_FILE %q: %w", caCertFile, err)}
|
||||
}
|
||||
|
||||
rootCAs, err := x509.SystemCertPool()
|
||||
if err != nil {
|
||||
return nil, &caCertFileError{fmt.Errorf("load system CA certificate pool: %w", err)}
|
||||
}
|
||||
if !rootCAs.AppendCertsFromPEM(caCertPEM) {
|
||||
return nil, &caCertFileError{fmt.Errorf("CA_CERT_FILE %q does not contain any valid PEM certificates", caCertFile)}
|
||||
}
|
||||
|
||||
return &tls.Config{RootCAs: rootCAs}, nil
|
||||
}
|
||||
|
||||
// getOptions returns the WebSocket client options, creating them if necessary.
|
||||
// It configures the connection URL, TLS settings, and authentication headers.
|
||||
func (client *WebSocketClient) getOptions() *gws.ClientOption {
|
||||
@@ -132,7 +175,7 @@ func (client *WebSocketClient) getOptions() *gws.ClientOption {
|
||||
|
||||
client.options = &gws.ClientOption{
|
||||
Addr: client.hubURL.String(),
|
||||
TlsConfig: &tls.Config{InsecureSkipVerify: true},
|
||||
TlsConfig: client.tlsConfig,
|
||||
RequestHeader: http.Header{
|
||||
"User-Agent": []string{getUserAgent()},
|
||||
"X-Token": []string{client.token},
|
||||
|
||||
@@ -4,6 +4,16 @@ package agent
|
||||
|
||||
import (
|
||||
"crypto/ed25519"
|
||||
"crypto/rand"
|
||||
"crypto/rsa"
|
||||
"crypto/tls"
|
||||
"crypto/x509"
|
||||
"crypto/x509/pkix"
|
||||
"encoding/pem"
|
||||
"math/big"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
@@ -16,6 +26,7 @@ import (
|
||||
"github.com/henrygd/beszel/internal/common"
|
||||
|
||||
"github.com/fxamacker/cbor/v2"
|
||||
"github.com/lxzan/gws"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"golang.org/x/crypto/ssh"
|
||||
@@ -165,6 +176,155 @@ func TestWebSocketClient_GetOptions(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestWebSocketClient_TLSVerification(t *testing.T) {
|
||||
agent := createTestAgent(t)
|
||||
serverCert, serverCertPEM := newSelfSignedServerCertificate(t)
|
||||
upgrader := gws.NewUpgrader(&gws.BuiltinEventHandler{}, nil)
|
||||
server := httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
conn, err := upgrader.Upgrade(w, r)
|
||||
if err == nil {
|
||||
go conn.ReadLoop()
|
||||
}
|
||||
}))
|
||||
server.TLS = &tls.Config{Certificates: []tls.Certificate{serverCert}}
|
||||
server.StartTLS()
|
||||
t.Cleanup(server.Close)
|
||||
|
||||
caCertFile := filepath.Join(t.TempDir(), "hub-ca.crt")
|
||||
require.NoError(t, os.WriteFile(caCertFile, serverCertPEM, 0600))
|
||||
|
||||
newClient := func(t *testing.T, caCertFile string) *WebSocketClient {
|
||||
t.Helper()
|
||||
t.Setenv("BESZEL_AGENT_HUB_URL", server.URL)
|
||||
t.Setenv("BESZEL_AGENT_TOKEN", "test-token")
|
||||
t.Setenv("BESZEL_AGENT_CA_CERT_FILE", caCertFile)
|
||||
client, err := newWebSocketClient(agent)
|
||||
require.NoError(t, err)
|
||||
return client
|
||||
}
|
||||
|
||||
t.Run("system roots are used by default", func(t *testing.T) {
|
||||
client := newClient(t, "")
|
||||
assert.Nil(t, client.getOptions().TlsConfig)
|
||||
_, _, err := gws.NewClient(&gws.BuiltinEventHandler{}, client.getOptions())
|
||||
require.Error(t, err)
|
||||
})
|
||||
|
||||
t.Run("custom CA trusts self-signed certificate", func(t *testing.T) {
|
||||
systemRoots, err := x509.SystemCertPool()
|
||||
require.NoError(t, err)
|
||||
client := newClient(t, caCertFile)
|
||||
assert.Greater(t, len(client.getOptions().TlsConfig.RootCAs.Subjects()), len(systemRoots.Subjects()))
|
||||
conn, _, err := gws.NewClient(&gws.BuiltinEventHandler{}, client.getOptions())
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, conn.NetConn().Close())
|
||||
})
|
||||
|
||||
t.Run("custom CA does not bypass hostname verification", func(t *testing.T) {
|
||||
client := newClient(t, caCertFile)
|
||||
client.getOptions().TlsConfig.ServerName = "wrong.example.com"
|
||||
_, _, err := gws.NewClient(&gws.BuiltinEventHandler{}, client.getOptions())
|
||||
require.Error(t, err)
|
||||
})
|
||||
}
|
||||
|
||||
func TestWebSocketClient_NonTLSConnection(t *testing.T) {
|
||||
agent := createTestAgent(t)
|
||||
upgrader := gws.NewUpgrader(&gws.BuiltinEventHandler{}, nil)
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
conn, err := upgrader.Upgrade(w, r)
|
||||
if err == nil {
|
||||
go conn.ReadLoop()
|
||||
}
|
||||
}))
|
||||
t.Cleanup(server.Close)
|
||||
|
||||
t.Setenv("BESZEL_AGENT_HUB_URL", server.URL)
|
||||
t.Setenv("BESZEL_AGENT_TOKEN", "test-token")
|
||||
t.Setenv("BESZEL_AGENT_CA_CERT_FILE", "")
|
||||
client, err := newWebSocketClient(agent)
|
||||
require.NoError(t, err)
|
||||
assert.Nil(t, client.getOptions().TlsConfig)
|
||||
|
||||
conn, _, err := gws.NewClient(&gws.BuiltinEventHandler{}, client.getOptions())
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, conn.NetConn().Close())
|
||||
}
|
||||
|
||||
func TestGetTLSConfigErrors(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
testCases := []struct {
|
||||
name string
|
||||
path string
|
||||
contents []byte
|
||||
errorMatch string
|
||||
}{
|
||||
{
|
||||
name: "missing file",
|
||||
path: filepath.Join(tempDir, "missing.pem"),
|
||||
errorMatch: "read CA_CERT_FILE",
|
||||
},
|
||||
{
|
||||
name: "unreadable path",
|
||||
path: tempDir,
|
||||
errorMatch: "read CA_CERT_FILE",
|
||||
},
|
||||
{
|
||||
name: "empty file",
|
||||
path: filepath.Join(tempDir, "empty.pem"),
|
||||
contents: []byte{},
|
||||
errorMatch: "does not contain any valid PEM certificates",
|
||||
},
|
||||
{
|
||||
name: "malformed file",
|
||||
path: filepath.Join(tempDir, "malformed.pem"),
|
||||
contents: []byte("not a PEM certificate"),
|
||||
errorMatch: "does not contain any valid PEM certificates",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range testCases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if tc.contents != nil {
|
||||
require.NoError(t, os.WriteFile(tc.path, tc.contents, 0600))
|
||||
}
|
||||
t.Setenv("BESZEL_AGENT_CA_CERT_FILE", tc.path)
|
||||
|
||||
tlsConfig, err := getTLSConfig()
|
||||
require.Error(t, err)
|
||||
assert.Nil(t, tlsConfig)
|
||||
assert.Contains(t, err.Error(), tc.errorMatch)
|
||||
assert.Contains(t, err.Error(), tc.path)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func newSelfSignedServerCertificate(t *testing.T) (tls.Certificate, []byte) {
|
||||
t.Helper()
|
||||
privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
|
||||
require.NoError(t, err)
|
||||
|
||||
template := &x509.Certificate{
|
||||
SerialNumber: big.NewInt(1),
|
||||
Subject: pkix.Name{CommonName: "127.0.0.1"},
|
||||
NotBefore: time.Now().Add(-time.Hour),
|
||||
NotAfter: time.Now().Add(time.Hour),
|
||||
IPAddresses: []net.IP{net.ParseIP("127.0.0.1")},
|
||||
KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageKeyEncipherment | x509.KeyUsageCertSign,
|
||||
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
|
||||
BasicConstraintsValid: true,
|
||||
IsCA: true,
|
||||
}
|
||||
certDER, err := x509.CreateCertificate(rand.Reader, template, template, &privateKey.PublicKey, privateKey)
|
||||
require.NoError(t, err)
|
||||
|
||||
certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certDER})
|
||||
keyPEM := pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(privateKey)})
|
||||
certificate, err := tls.X509KeyPair(certPEM, keyPEM)
|
||||
require.NoError(t, err)
|
||||
return certificate, certPEM
|
||||
}
|
||||
|
||||
// TestWebSocketClient_VerifySignature tests signature verification
|
||||
func TestWebSocketClient_VerifySignature(t *testing.T) {
|
||||
agent := createTestAgent(t)
|
||||
|
||||
@@ -87,6 +87,10 @@ func (c *ConnectionManager) Start(serverOptions ServerOptions) error {
|
||||
|
||||
wsClient, err := newWebSocketClient(c.agent)
|
||||
if err != nil {
|
||||
var caCertErr *caCertFileError
|
||||
if errors.As(err, &caCertErr) {
|
||||
return err
|
||||
}
|
||||
slog.Warn("Error creating WebSocket client", "err", err)
|
||||
}
|
||||
c.wsClient = wsClient
|
||||
|
||||
@@ -265,6 +265,19 @@ func TestConnectionManager_StartWithInvalidConfig(t *testing.T) {
|
||||
assert.Error(t, err, "Should error when starting already started connection manager")
|
||||
}
|
||||
|
||||
func TestConnectionManager_StartRejectsInvalidCACertFile(t *testing.T) {
|
||||
agent := createTestAgent(t)
|
||||
cm := agent.connectionManager
|
||||
t.Setenv("BESZEL_AGENT_HUB_URL", "https://hub.example.com")
|
||||
t.Setenv("BESZEL_AGENT_TOKEN", "test-token")
|
||||
t.Setenv("BESZEL_AGENT_CA_CERT_FILE", t.TempDir())
|
||||
|
||||
err := cm.Start(ServerOptions{})
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "read CA_CERT_FILE")
|
||||
assert.Nil(t, cm.eventChan)
|
||||
}
|
||||
|
||||
// TestConnectionManager_CloseWebSocket tests WebSocket closing
|
||||
func TestConnectionManager_CloseWebSocket(t *testing.T) {
|
||||
agent := createTestAgent(t)
|
||||
|
||||
Reference in New Issue
Block a user