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 doSelfUpdate = selfUpdate
|
||||||
|
|
||||||
var errRestartRequested = fmt.Errorf("restart requested")
|
var errRestartRequested = fmt.Errorf("restart requested")
|
||||||
|
var joinTimeout = 10 * time.Second
|
||||||
|
|
||||||
const maxJobPayloadBytes = 4 << 20 // 4 MB — well above any legitimate job list
|
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)
|
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
|
// Join channel
|
||||||
joinPayload, _ := json.Marshal(map[string]string{"token": token})
|
joinPayload, _ := json.Marshal(map[string]string{"token": token})
|
||||||
joinMsg := channelMsg{
|
joinMsg := channelMsg{
|
||||||
|
|
@ -166,6 +181,29 @@ func runSession(ctx context.Context, baseURL, token string) error {
|
||||||
}
|
}
|
||||||
slog.Debug("sent channel join request")
|
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
|
// Writer goroutine - serializes all writes to the WebSocket
|
||||||
var writerWg sync.WaitGroup
|
var writerWg sync.WaitGroup
|
||||||
writeErrCh := make(chan error, 1)
|
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)
|
heartbeatTicker := time.NewTicker(60 * time.Second)
|
||||||
defer heartbeatTicker.Stop()
|
defer heartbeatTicker.Stop()
|
||||||
channelHeartbeatTicker := time.NewTicker(25 * time.Second)
|
channelHeartbeatTicker := time.NewTicker(25 * time.Second)
|
||||||
|
|
|
||||||
137
agent_test.go
137
agent_test.go
|
|
@ -5,6 +5,8 @@ import (
|
||||||
"encoding/base64"
|
"encoding/base64"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"net"
|
||||||
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"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