security: validate join reply before entering main event loop

Previously the agent proceeded to the main loop without confirming
the server accepted the channel join. A rejected join (e.g. invalid
token) left the agent running but silently non-functional. Now waits
for the join reply with a 10s timeout, returns error on rejection.

Reader goroutine starts before join send so the reply can be received.
Writer goroutine starts after join confirmation.
This commit is contained in:
Graham McIntire 2026-02-12 11:02:38 -06:00
parent 13e923d251
commit 0f8e8fc183
No known key found for this signature in database
2 changed files with 175 additions and 14 deletions

View file

@ -24,6 +24,7 @@ var osExit = os.Exit
var doSelfUpdate = selfUpdate
var errRestartRequested = fmt.Errorf("restart requested")
var joinTimeout = 10 * time.Second
const maxJobPayloadBytes = 4 << 20 // 4 MB — well above any legitimate job list
@ -152,6 +153,20 @@ func runSession(ctx context.Context, baseURL, token string) error {
sendMsg(event, payload)
}
// Reader goroutine — must start before join so we can receive the reply
msgCh := make(chan []byte, 100)
errCh := make(chan error, 1)
go func() {
for {
data, _, err := ws.ReadMessage()
if err != nil {
errCh <- err
return
}
msgCh <- data
}
}()
// Join channel
joinPayload, _ := json.Marshal(map[string]string{"token": token})
joinMsg := channelMsg{
@ -166,6 +181,29 @@ func runSession(ctx context.Context, baseURL, token string) error {
}
slog.Debug("sent channel join request")
// Wait for join reply before entering main loop
select {
case data := <-msgCh:
var reply channelMsg
if err := json.Unmarshal(data, &reply); err != nil {
return fmt.Errorf("join reply unmarshal: %w", err)
}
if reply.Event == "phx_reply" {
var status struct {
Status string `json:"status"`
Response any `json:"response"`
}
if err := json.Unmarshal(reply.Payload, &status); err == nil && status.Status != "ok" {
return fmt.Errorf("join rejected: %s", status.Status)
}
}
slog.Info("channel joined")
case err := <-errCh:
return fmt.Errorf("read during join: %w", err)
case <-time.After(joinTimeout):
return fmt.Errorf("join timeout")
}
// Writer goroutine - serializes all writes to the WebSocket
var writerWg sync.WaitGroup
writeErrCh := make(chan error, 1)
@ -184,20 +222,6 @@ func runSession(ctx context.Context, baseURL, token string) error {
}
}()
// Reader goroutine - reads messages and dispatches
msgCh := make(chan []byte, 100)
errCh := make(chan error, 1)
go func() {
for {
data, _, err := ws.ReadMessage()
if err != nil {
errCh <- err
return
}
msgCh <- data
}
}()
heartbeatTicker := time.NewTicker(60 * time.Second)
defer heartbeatTicker.Stop()
channelHeartbeatTicker := time.NewTicker(25 * time.Second)

View file

@ -5,6 +5,8 @@ import (
"encoding/base64"
"encoding/json"
"fmt"
"net"
"strings"
"sync"
"testing"
"time"
@ -542,3 +544,138 @@ func TestDispatchJob(t *testing.T) {
}
})
}
func TestRunSessionRejectsFailedJoin(t *testing.T) {
origTimeout := joinTimeout
defer func() { joinTimeout = origTimeout }()
joinTimeout = 2 * time.Second
// Start a fake WebSocket server that accepts the upgrade then sends a phx_error join reply
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() }()
// Read HTTP upgrade request
buf := make([]byte, 4096)
n, _ := conn.Read(buf)
reqStr := string(buf[:n])
// Extract key and compute accept
key := extractWSKey(reqStr)
accept := computeAcceptKey(key)
// Send valid 101 upgrade
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))
// Read the join message (masked WebSocket frame) — just consume it
frameBuf := make([]byte, 4096)
_, _ = conn.Read(frameBuf)
// Send phx_error reply as an unmasked text frame
reply, _ := json.Marshal(channelMsg{
Topic: "agent:agent-0",
Event: "phx_reply",
Payload: json.RawMessage(`{"status":"error","response":{"reason":"invalid token"}}`),
Ref: strPtr("1"),
})
frame := makeTextFrame(reply)
_, _ = conn.Write(frame)
// Keep connection open for a bit
time.Sleep(time.Second)
}()
addr := ln.Addr().String()
err = runSession(context.Background(), "ws://"+addr, "bad-token")
if err == nil {
t.Fatal("expected error from rejected join")
}
if !strings.Contains(err.Error(), "join rejected") {
t.Errorf("expected 'join rejected' in error, got: %v", err)
}
}
func TestRunSessionJoinTimeout(t *testing.T) {
origTimeout := joinTimeout
defer func() { joinTimeout = origTimeout }()
joinTimeout = 500 * time.Millisecond
// Server that upgrades but never sends a join reply
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)
n, _ := conn.Read(buf)
reqStr := string(buf[:n])
key := extractWSKey(reqStr)
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))
// Read join frame but never reply
frameBuf := make([]byte, 4096)
_, _ = conn.Read(frameBuf)
time.Sleep(5 * time.Second)
}()
addr := ln.Addr().String()
start := time.Now()
err = runSession(context.Background(), "ws://"+addr, "token")
elapsed := time.Since(start)
if err == nil {
t.Fatal("expected timeout error")
}
if !strings.Contains(err.Error(), "join timeout") {
t.Errorf("expected 'join timeout' in error, got: %v", err)
}
if elapsed > 3*time.Second {
t.Errorf("took too long (%v), timeout didn't trigger", elapsed)
}
}
// extractWSKey extracts the Sec-WebSocket-Key from a raw HTTP request.
func extractWSKey(req string) string {
for _, line := range strings.Split(req, "\r\n") {
lower := strings.ToLower(line)
if strings.HasPrefix(lower, "sec-websocket-key: ") {
return strings.TrimSpace(line[len("Sec-WebSocket-Key: "):])
}
}
return ""
}
// makeTextFrame creates an unmasked WebSocket text frame (server→client).
func makeTextFrame(payload []byte) []byte {
length := len(payload)
var frame []byte
frame = append(frame, 0x81) // FIN + text
if length <= 125 {
frame = append(frame, byte(length))
} else if length <= 65535 {
frame = append(frame, 126, byte(length>>8), byte(length))
}
frame = append(frame, payload...)
return frame
}