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:
parent
13e923d251
commit
0f8e8fc183
2 changed files with 175 additions and 14 deletions
52
agent.go
52
agent.go
|
|
@ -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)
|
||||
|
|
|
|||
137
agent_test.go
137
agent_test.go
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue