From aa7ca96e95b11f129969753e2bf1e781531a3aba Mon Sep 17 00:00:00 2001 From: Graham McIntire Date: Thu, 12 Feb 2026 10:33:17 -0600 Subject: [PATCH] harden security: file permissions, WebSocket accept validation, backoff jitter - Self-update: write binary with 0700 permissions instead of 0755 - WebSocket: validate Sec-WebSocket-Accept per RFC 6455 to prevent MITM - WebSocket: add computeAcceptKey helper with correct RFC 6455 GUID - Agent: add nextBackoff with 25% jitter to prevent thundering herd --- agent.go | 13 +++++++++- agent_test.go | 32 ++++++++++++++++++++++++ update.go | 2 +- update_test.go | 39 +++++++++++++++++++++++++++++ websocket.go | 30 ++++++++++++++++++++++ websocket_test.go | 64 ++++++++++++++++++++++++++++++++++++++++++++--- 6 files changed, 174 insertions(+), 6 deletions(-) diff --git a/agent.go b/agent.go index 633d76f..09fbd06 100644 --- a/agent.go +++ b/agent.go @@ -6,6 +6,7 @@ import ( "encoding/json" "fmt" "log/slog" + "math/rand/v2" "os" "runtime" "strconv" @@ -56,7 +57,7 @@ func runAgent(ctx context.Context, wsURL, token string) { return case <-time.After(retryDelay): } - retryDelay = min(retryDelay*2, maxRetry) + retryDelay = nextBackoff(retryDelay, maxRetry) } } @@ -363,4 +364,14 @@ func dispatchJob( } } +// nextBackoff doubles the current delay (capped at max) and adds up to 25% jitter. +func nextBackoff(current, maxDelay time.Duration) time.Duration { + next := current * 2 + if next > maxDelay { + next = maxDelay + } + jitter := time.Duration(rand.Int64N(int64(next / 4))) + return next + jitter +} + func strPtr(s string) *string { return &s } diff --git a/agent_test.go b/agent_test.go index 12163f5..a9fe0d4 100644 --- a/agent_test.go +++ b/agent_test.go @@ -336,6 +336,38 @@ func TestHandleMessage(t *testing.T) { }) } +func TestNextBackoff(t *testing.T) { + maxDelay := 60 * time.Second + + // Test doubling with jitter + for i := 0; i < 100; i++ { + current := 2 * time.Second + next := nextBackoff(current, maxDelay) + doubled := current * 2 + maxWithJitter := doubled + doubled/4 + if next < doubled || next > maxWithJitter { + t.Errorf("nextBackoff(%v) = %v, want in [%v, %v]", current, next, doubled, maxWithJitter) + } + } + + // Test cap at max + for i := 0; i < 100; i++ { + next := nextBackoff(30*time.Second, maxDelay) + if next > maxDelay+maxDelay/4 { + t.Errorf("nextBackoff(30s) = %v, exceeded max+jitter", next) + } + } + + // Test that already-at-max stays at max (with jitter) + for i := 0; i < 100; i++ { + next := nextBackoff(maxDelay, maxDelay) + maxWithJitter := maxDelay + maxDelay/4 + if next < maxDelay || next > maxWithJitter { + t.Errorf("nextBackoff(max) = %v, want in [%v, %v]", next, maxDelay, maxWithJitter) + } + } +} + func TestDispatchJob(t *testing.T) { t.Run("MIKROTIK", func(t *testing.T) { origDial := mikrotikDial diff --git a/update.go b/update.go index 77dcc38..db936cc 100644 --- a/update.go +++ b/update.go @@ -66,7 +66,7 @@ func selfUpdate(downloadURL, expectedChecksum string) error { } tempPath := currentExe + ".update" - if err := osWriteFile(tempPath, body, 0755); err != nil { + if err := osWriteFile(tempPath, body, 0700); err != nil { return fmt.Errorf("write temp: %w", err) } diff --git a/update_test.go b/update_test.go index 373ba11..091c5f2 100644 --- a/update_test.go +++ b/update_test.go @@ -218,6 +218,45 @@ func TestSelfUpdateChecksumMatch(t *testing.T) { } } +func TestSelfUpdateFilePermissions(t *testing.T) { + origExe := osExecutable + origWrite := osWriteFile + origRename := osRename + defer func() { + osExecutable = origExe + osWriteFile = origWrite + osRename = origRename + }() + osExecutable = func() (string, error) { return "/tmp/test-agent", nil } + osRename = func(oldpath, newpath string) error { + return fmt.Errorf("stop here") // stop before re-exec + } + + var capturedPerm os.FileMode + osWriteFile = func(name string, data []byte, perm os.FileMode) error { + capturedPerm = perm + return nil + } + + body := []byte("binary data") + checksum := fmt.Sprintf("%x", sha256.Sum256(body)) + + srv := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _, _ = w.Write(body) + })) + defer srv.Close() + + origGet := httpGet + defer func() { httpGet = origGet }() + httpGet = srv.Client().Get + + _ = selfUpdate(rewriteToHTTPS(srv.URL), checksum) + + if capturedPerm != 0700 { + t.Errorf("expected file permissions 0700, got %o", capturedPerm) + } +} + func TestSelfUpdateTooLarge(t *testing.T) { origMax := maxUpdateSize defer func() { maxUpdateSize = origMax }() diff --git a/websocket.go b/websocket.go index 2bca9f2..20d21e5 100644 --- a/websocket.go +++ b/websocket.go @@ -3,6 +3,7 @@ package main import ( "bufio" "crypto/rand" + "crypto/sha1" "crypto/tls" "encoding/base64" "encoding/binary" @@ -14,6 +15,16 @@ import ( "sync" ) +// websocketGUID is the magic GUID from RFC 6455 Section 4.2.2. +const websocketGUID = "258EAFA5-E914-47DA-95CA-C5AB0DC85B11" + +// computeAcceptKey computes the expected Sec-WebSocket-Accept value per RFC 6455. +func computeAcceptKey(key string) string { + h := sha1.New() + h.Write([]byte(key + websocketGUID)) + return base64.StdEncoding.EncodeToString(h.Sum(nil)) +} + const ( opText = 1 opBinary = 2 @@ -94,6 +105,25 @@ func WSDial(rawURL string) (*WSConn, error) { return nil, fmt.Errorf("handshake failed: %s", strings.SplitN(resp, "\r\n", 2)[0]) } + // Verify Sec-WebSocket-Accept per RFC 6455 + expectedAccept := computeAcceptKey(key) + acceptFound := false + for _, line := range strings.Split(resp, "\r\n") { + if strings.HasPrefix(strings.ToLower(line), "sec-websocket-accept: ") { + actual := strings.TrimSpace(line[len("Sec-WebSocket-Accept: "):]) + if actual != expectedAccept { + _ = conn.Close() + return nil, fmt.Errorf("invalid accept key: got %q, want %q", actual, expectedAccept) + } + acceptFound = true + break + } + } + if !acceptFound { + _ = conn.Close() + return nil, fmt.Errorf("missing Sec-WebSocket-Accept header") + } + return &WSConn{conn: conn, reader: bufio.NewReaderSize(conn, 8192)}, nil } diff --git a/websocket_test.go b/websocket_test.go index cdd364d..9543286 100644 --- a/websocket_test.go +++ b/websocket_test.go @@ -283,7 +283,6 @@ func TestWriteFrameHeaderError(t *testing.T) { } } - func TestReadMessageError(t *testing.T) { // Empty buffer causes immediate EOF on readFrame buf := bytes.NewBuffer(nil) @@ -422,9 +421,13 @@ func TestWSDial(t *testing.T) { buf := make([]byte, 4096) n, _ := conn.Read(buf) - _ = string(buf[:n]) // Read the request + reqStr := string(buf[:n]) - resp := "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\n\r\n" + // Extract Sec-WebSocket-Key from request + key := extractHeader(reqStr, "Sec-WebSocket-Key") + accept := computeAcceptKey(key) + + resp := "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\nSec-WebSocket-Accept: " + accept + "\r\n\r\n" _, _ = conn.Write([]byte(resp)) // Keep connection open briefly for the test @@ -578,9 +581,11 @@ func TestWSDialRealHTTPServer(t *testing.T) { http.Error(w, "hijack not supported", 500) return } + key := r.Header.Get("Sec-WebSocket-Key") + accept := computeAcceptKey(key) conn, brw, _ := hj.Hijack() defer func() { _ = conn.Close() }() - _, _ = brw.WriteString("HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\n\r\n") + _, _ = brw.WriteString("HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\nSec-WebSocket-Accept: " + accept + "\r\n\r\n") _ = brw.Flush() // Keep alive briefly buf := make([]byte, 1) @@ -598,6 +603,47 @@ func TestWSDialRealHTTPServer(t *testing.T) { _ = ws.Close() } +func TestComputeAcceptKey(t *testing.T) { + // RFC 6455 Section 4.2.2 test vector + got := computeAcceptKey("dGhlIHNhbXBsZSBub25jZQ==") + want := "s3pPLMBiTxaQ9kYGzzhZRbK+xOo=" + if got != want { + t.Errorf("computeAcceptKey = %q, want %q", got, want) + } +} + +func TestWSDialRejectsInvalidAcceptKey(t *testing.T) { + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer func() { _ = ln.Close() }() + + go func() { + conn, err := ln.Accept() + if err != nil { + return + } + defer func() { _ = conn.Close() }() + + buf := make([]byte, 4096) + _, _ = conn.Read(buf) + + // Send 101 with a wrong accept key + resp := "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\nSec-WebSocket-Accept: INVALID_KEY\r\n\r\n" + _, _ = conn.Write([]byte(resp)) + }() + + addr := ln.Addr().String() + _, err = WSDial("ws://" + addr + "/socket") + if err == nil { + t.Error("expected error for invalid accept key") + } + if !strings.Contains(err.Error(), "accept key") { + t.Errorf("expected 'accept key' in error, got: %v", err) + } +} + func testWSConn(rw io.ReadWriteCloser) *WSConn { return &WSConn{conn: rw, reader: bufio.NewReader(rw)} } @@ -614,6 +660,16 @@ func (n *nopCloser) Read(p []byte) (int, error) { return n.readWriter.Read(p) } func (n *nopCloser) Write(p []byte) (int, error) { return n.readWriter.Write(p) } func (n *nopCloser) Close() error { return nil } +// extractHeader extracts a header value from a raw HTTP request string. +func extractHeader(req, name string) string { + for _, line := range strings.Split(req, "\r\n") { + if strings.HasPrefix(strings.ToLower(line), strings.ToLower(name)+": ") { + return strings.TrimSpace(line[len(name)+2:]) + } + } + return "" +} + // captureWriter captures written data while reading from a separate Reader. type captureWriter struct { Reader *bytes.Buffer