diff --git a/agent.go b/agent.go index 8337ab6..6058a0f 100644 --- a/agent.go +++ b/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) diff --git a/agent_test.go b/agent_test.go index 93542ea..fe2b68b 100644 --- a/agent_test.go +++ b/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 +}