From 26e80bbe17a8b9e7686e76f81566ab6cdd2cb82d Mon Sep 17 00:00:00 2001 From: Alex Nahas Date: Sat, 3 Oct 2026 12:28:37 -0700 Subject: [PATCH] Require an exit code before exec reports success Non-interactive exec returned 0 when the WebSocket closed without an exit code message, and the exit code could lose a select race to the reader's done channel. Treat a missing or malformed exit code as an error. --- pkg/cmd/exec.go | 17 ++------- pkg/cmd/exec_test.go | 88 ++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 92 insertions(+), 13 deletions(-) create mode 100644 pkg/cmd/exec_test.go diff --git a/pkg/cmd/exec.go b/pkg/cmd/exec.go index d53efe5..04fb3e3 100644 --- a/pkg/cmd/exec.go +++ b/pkg/cmd/exec.go @@ -295,7 +295,6 @@ func runExecInteractive(ws *websocket.Conn) (int, error) { func runExecNonInteractive(ws *websocket.Conn) (int, error) { errCh := make(chan error, 2) exitCodeCh := make(chan int, 1) - doneCh := make(chan struct{}) // Forward stdin to WebSocket go func() { @@ -319,26 +318,20 @@ func runExecNonInteractive(ws *websocket.Conn) (int, error) { // Forward WebSocket to stdout go func() { - defer close(doneCh) for { msgType, message, err := ws.ReadMessage() if err != nil { - if websocket.IsCloseError(err, websocket.CloseNormalClosure, websocket.CloseGoingAway, websocket.CloseAbnormalClosure) || - err == io.EOF { - exitCodeCh <- 0 - return - } - errCh <- fmt.Errorf("websocket read error: %w", err) + errCh <- fmt.Errorf("websocket closed before receiving an exit code: %w", err) return } // Check for exit code message if msgType == websocket.TextMessage && bytes.Contains(message, []byte("exitCode")) { var exitMsg struct { - ExitCode int `json:"exitCode"` + ExitCode *int `json:"exitCode"` } - if json.Unmarshal(message, &exitMsg) == nil { - exitCodeCh <- exitMsg.ExitCode + if json.Unmarshal(message, &exitMsg) == nil && exitMsg.ExitCode != nil { + exitCodeCh <- *exitMsg.ExitCode return } } @@ -355,7 +348,5 @@ func runExecNonInteractive(ws *websocket.Conn) (int, error) { return 255, err case exitCode := <-exitCodeCh: return exitCode, nil - case <-doneCh: - return 0, nil } } diff --git a/pkg/cmd/exec_test.go b/pkg/cmd/exec_test.go new file mode 100644 index 0000000..736edb3 --- /dev/null +++ b/pkg/cmd/exec_test.go @@ -0,0 +1,88 @@ +package cmd + +import ( + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/gorilla/websocket" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// dialExecTestServer starts a WebSocket server that runs serve on the +// accepted connection and returns a client connected to it. +func dialExecTestServer(t *testing.T, serve func(*websocket.Conn)) *websocket.Conn { + t.Helper() + upgrader := websocket.Upgrader{} + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + ws, err := upgrader.Upgrade(w, r, nil) + if err != nil { + return + } + defer ws.Close() + serve(ws) + })) + t.Cleanup(server.Close) + + ws, _, err := websocket.DefaultDialer.Dial("ws"+strings.TrimPrefix(server.URL, "http"), nil) + require.NoError(t, err) + t.Cleanup(func() { ws.Close() }) + return ws +} + +func TestRunExecNonInteractivePropagatesExitCode(t *testing.T) { + ws := dialExecTestServer(t, func(ws *websocket.Conn) { + _ = ws.WriteMessage(websocket.TextMessage, []byte(`{"exitCode":3}`)) + _ = ws.WriteMessage(websocket.CloseMessage, websocket.FormatCloseMessage(websocket.CloseNormalClosure, "")) + }) + + exitCode, err := runExecNonInteractive(ws) + require.NoError(t, err) + assert.Equal(t, 3, exitCode) +} + +func TestRunExecNonInteractiveRequiresExitCode(t *testing.T) { + tests := []struct { + name string + serve func(*websocket.Conn) + }{ + { + name: "normal close without exit code", + serve: func(ws *websocket.Conn) { + _ = ws.WriteMessage(websocket.CloseMessage, websocket.FormatCloseMessage(websocket.CloseNormalClosure, "")) + }, + }, + { + name: "malformed exit code", + serve: func(ws *websocket.Conn) { + _ = ws.WriteMessage(websocket.TextMessage, []byte(`{"exitCode":"oops"}`)) + _ = ws.WriteMessage(websocket.CloseMessage, websocket.FormatCloseMessage(websocket.CloseNormalClosure, "")) + }, + }, + { + name: "null exit code", + serve: func(ws *websocket.Conn) { + _ = ws.WriteMessage(websocket.TextMessage, []byte(`{"exitCode":null}`)) + _ = ws.WriteMessage(websocket.CloseMessage, websocket.FormatCloseMessage(websocket.CloseNormalClosure, "")) + }, + }, + { + name: "connection dropped before exit code", + serve: func(ws *websocket.Conn) { + _ = ws.NetConn().Close() + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ws := dialExecTestServer(t, tt.serve) + + exitCode, err := runExecNonInteractive(ws) + require.ErrorContains(t, err, "before receiving an exit code") + assert.NotEqual(t, 0, exitCode) + }) + } +}