From f104f31ee35ed1f45129d1bd7ad3df9e26dd2f88 Mon Sep 17 00:00:00 2001 From: hank Date: Wed, 2 Sep 2026 13:55:54 -0400 Subject: [PATCH] Merge commit from fork --- agent/client.go | 45 ++++++++- agent/client_test.go | 160 +++++++++++++++++++++++++++++++ agent/connection_manager.go | 4 + agent/connection_manager_test.go | 13 +++ 4 files changed, 221 insertions(+), 1 deletion(-) diff --git a/agent/client.go b/agent/client.go index 2436b294..0dc80fa6 100644 --- a/agent/client.go +++ b/agent/client.go @@ -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}, diff --git a/agent/client_test.go b/agent/client_test.go index 1b987eef..a4598fc1 100644 --- a/agent/client_test.go +++ b/agent/client_test.go @@ -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) diff --git a/agent/connection_manager.go b/agent/connection_manager.go index 30b5c543..f9854f2b 100644 --- a/agent/connection_manager.go +++ b/agent/connection_manager.go @@ -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 diff --git a/agent/connection_manager_test.go b/agent/connection_manager_test.go index 8aba1579..b78fc5fc 100644 --- a/agent/connection_manager_test.go +++ b/agent/connection_manager_test.go @@ -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)