Merge commit from fork

This commit is contained in:
hank
2026-09-02 13:55:54 -04:00
committed by GitHub
parent a1ca51608a
commit f104f31ee3
4 changed files with 221 additions and 1 deletions
+44 -1
View File
@@ -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},
+160
View File
@@ -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)
+4
View File
@@ -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
+13
View File
@@ -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)