diff --git a/agent/agentcontainers/api.go b/agent/agentcontainers/api.go index 8e056fa666..e69e3f2d99 100644 --- a/agent/agentcontainers/api.go +++ b/agent/agentcontainers/api.go @@ -779,10 +779,13 @@ func (api *API) watchContainers(rw http.ResponseWriter, r *http.Request) { // close frames. _ = conn.CloseRead(context.Background()) + ctx, cancel := context.WithCancel(ctx) + defer cancel() + ctx, wsNetConn := codersdk.WebsocketNetConn(ctx, conn, websocket.MessageText) defer wsNetConn.Close() - go httpapi.Heartbeat(ctx, conn) + go httpapi.HeartbeatClose(ctx, api.logger, cancel, conn) updateCh := make(chan struct{}, 1) diff --git a/coderd/httpapi/httpapi.go b/coderd/httpapi/httpapi.go index bccc58ab37..c2dd81bbc3 100644 --- a/coderd/httpapi/httpapi.go +++ b/coderd/httpapi/httpapi.go @@ -16,6 +16,7 @@ import ( "github.com/go-playground/validator/v10" "golang.org/x/xerrors" + "cdr.dev/slog/v3" "github.com/coder/coder/v2/coderd/httpapi/httpapiconstraints" "github.com/coder/coder/v2/coderd/tracing" "github.com/coder/coder/v2/codersdk" @@ -418,79 +419,85 @@ func ServerSentEventSender(rw http.ResponseWriter, r *http.Request) ( // open a workspace in multiple tabs, the entire UI can start to lock up. // WebSockets have no such limitation, no matter what HTTP protocol was used to // establish the connection. -func OneWayWebSocketEventSender(rw http.ResponseWriter, r *http.Request) ( +func OneWayWebSocketEventSender(log slog.Logger) func(rw http.ResponseWriter, r *http.Request) ( func(event codersdk.ServerSentEvent) error, <-chan struct{}, error, ) { - ctx, cancel := context.WithCancel(r.Context()) - r = r.WithContext(ctx) - socket, err := websocket.Accept(rw, r, nil) - if err != nil { - cancel() - return nil, nil, xerrors.Errorf("cannot establish connection: %w", err) - } - go Heartbeat(ctx, socket) - - eventC := make(chan codersdk.ServerSentEvent) - socketErrC := make(chan websocket.CloseError, 1) - closed := make(chan struct{}) - go func() { - defer cancel() - defer close(closed) - - for { - select { - case event := <-eventC: - writeCtx, cancel := context.WithTimeout(ctx, 10*time.Second) - err := wsjson.Write(writeCtx, socket, event) - cancel() - if err == nil { - continue - } - _ = socket.Close(websocket.StatusInternalError, "Unable to send newest message") - case err := <-socketErrC: - _ = socket.Close(err.Code, err.Reason) - case <-ctx.Done(): - _ = socket.Close(websocket.StatusNormalClosure, "Connection closed") - } - return - } - }() - - // We have some tools in the UI code to help enforce one-way WebSocket - // connections, but there's still the possibility that the client could send - // a message when it's not supposed to. If that happens, the client likely - // forgot to use those tools, and communication probably can't be trusted. - // Better to just close the socket and force the UI to fix its mess - go func() { - _, _, err := socket.Read(ctx) - if errors.Is(err, context.Canceled) { - return - } + return func(rw http.ResponseWriter, r *http.Request) ( + func(event codersdk.ServerSentEvent) error, + <-chan struct{}, + error, + ) { + ctx, cancel := context.WithCancel(r.Context()) + r = r.WithContext(ctx) + socket, err := websocket.Accept(rw, r, nil) if err != nil { - socketErrC <- websocket.CloseError{ - Code: websocket.StatusInternalError, - Reason: "Unable to process invalid message from client", + cancel() + return nil, nil, xerrors.Errorf("cannot establish connection: %w", err) + } + go HeartbeatClose(ctx, log, cancel, socket) + + eventC := make(chan codersdk.ServerSentEvent) + socketErrC := make(chan websocket.CloseError, 1) + closed := make(chan struct{}) + go func() { + defer cancel() + defer close(closed) + + for { + select { + case event := <-eventC: + writeCtx, cancel := context.WithTimeout(ctx, 10*time.Second) + err := wsjson.Write(writeCtx, socket, event) + cancel() + if err == nil { + continue + } + _ = socket.Close(websocket.StatusInternalError, "Unable to send newest message") + case err := <-socketErrC: + _ = socket.Close(err.Code, err.Reason) + case <-ctx.Done(): + _ = socket.Close(websocket.StatusNormalClosure, "Connection closed") + } + return } - return - } - socketErrC <- websocket.CloseError{ - Code: websocket.StatusProtocolError, - Reason: "Clients cannot send messages for one-way WebSockets", - } - }() + }() - sendEvent := func(event codersdk.ServerSentEvent) error { - select { - case eventC <- event: - case <-ctx.Done(): - return ctx.Err() + // We have some tools in the UI code to help enforce one-way WebSocket + // connections, but there's still the possibility that the client could send + // a message when it's not supposed to. If that happens, the client likely + // forgot to use those tools, and communication probably can't be trusted. + // Better to just close the socket and force the UI to fix its mess + go func() { + _, _, err := socket.Read(ctx) + if errors.Is(err, context.Canceled) { + return + } + if err != nil { + socketErrC <- websocket.CloseError{ + Code: websocket.StatusInternalError, + Reason: "Unable to process invalid message from client", + } + return + } + socketErrC <- websocket.CloseError{ + Code: websocket.StatusProtocolError, + Reason: "Clients cannot send messages for one-way WebSockets", + } + }() + + sendEvent := func(event codersdk.ServerSentEvent) error { + select { + case eventC <- event: + case <-ctx.Done(): + return ctx.Err() + } + return nil } - return nil + + return sendEvent, closed, nil } - - return sendEvent, closed, nil } // WriteOAuth2Error writes an OAuth2-compliant error response per RFC 6749. diff --git a/coderd/httpapi/httpapi_test.go b/coderd/httpapi/httpapi_test.go index 44675e78a2..0fc6df8e8b 100644 --- a/coderd/httpapi/httpapi_test.go +++ b/coderd/httpapi/httpapi_test.go @@ -18,6 +18,7 @@ import ( "github.com/stretchr/testify/require" "golang.org/x/xerrors" + "cdr.dev/slog/v3/sloggers/slogtest" "github.com/coder/coder/v2/coderd/httpapi" "github.com/coder/coder/v2/codersdk" "github.com/coder/coder/v2/testutil" @@ -262,7 +263,7 @@ func TestOneWayWebSocketEventSender(t *testing.T) { req.Proto = p.proto writer := newOneWayWriter(t) - _, _, err := httpapi.OneWayWebSocketEventSender(writer, req) + _, _, err := httpapi.OneWayWebSocketEventSender(slogtest.Make(t, nil))(writer, req) require.ErrorContains(t, err, p.proto) } }) @@ -273,7 +274,7 @@ func TestOneWayWebSocketEventSender(t *testing.T) { ctx := testutil.Context(t, testutil.WaitShort) req := newBaseRequest(ctx) writer := newOneWayWriter(t) - send, _, err := httpapi.OneWayWebSocketEventSender(writer, req) + send, _, err := httpapi.OneWayWebSocketEventSender(slogtest.Make(t, nil))(writer, req) require.NoError(t, err) serverPayload := codersdk.ServerSentEvent{ @@ -299,7 +300,7 @@ func TestOneWayWebSocketEventSender(t *testing.T) { ctx, cancel := context.WithCancel(testutil.Context(t, testutil.WaitShort)) req := newBaseRequest(ctx) writer := newOneWayWriter(t) - _, done, err := httpapi.OneWayWebSocketEventSender(writer, req) + _, done, err := httpapi.OneWayWebSocketEventSender(slogtest.Make(t, nil))(writer, req) require.NoError(t, err) successC := make(chan bool) @@ -323,7 +324,7 @@ func TestOneWayWebSocketEventSender(t *testing.T) { ctx := testutil.Context(t, testutil.WaitShort) req := newBaseRequest(ctx) writer := newOneWayWriter(t) - _, done, err := httpapi.OneWayWebSocketEventSender(writer, req) + _, done, err := httpapi.OneWayWebSocketEventSender(slogtest.Make(t, nil))(writer, req) require.NoError(t, err) successC := make(chan bool) @@ -353,7 +354,7 @@ func TestOneWayWebSocketEventSender(t *testing.T) { ctx, cancel := context.WithCancel(testutil.Context(t, testutil.WaitShort)) req := newBaseRequest(ctx) writer := newOneWayWriter(t) - send, done, err := httpapi.OneWayWebSocketEventSender(writer, req) + send, done, err := httpapi.OneWayWebSocketEventSender(slogtest.Make(t, nil))(writer, req) require.NoError(t, err) successC := make(chan bool) @@ -394,7 +395,7 @@ func TestOneWayWebSocketEventSender(t *testing.T) { ctx := testutil.Context(t, timeout) req := newBaseRequest(ctx) writer := newOneWayWriter(t) - _, _, err := httpapi.OneWayWebSocketEventSender(writer, req) + _, _, err := httpapi.OneWayWebSocketEventSender(slogtest.Make(t, nil))(writer, req) require.NoError(t, err) type Result struct { diff --git a/coderd/httpapi/websocket.go b/coderd/httpapi/websocket.go index 397d7b94ab..ea9e505146 100644 --- a/coderd/httpapi/websocket.go +++ b/coderd/httpapi/websocket.go @@ -13,26 +13,7 @@ import ( const HeartbeatInterval time.Duration = 15 * time.Second -// Heartbeat loops to ping a WebSocket to keep it alive. -// Default idle connection timeouts are typically 60 seconds. -// See: https://docs.aws.amazon.com/elasticloadbalancing/latest/application/application-load-balancers.html#connection-idle-timeout -func Heartbeat(ctx context.Context, conn *websocket.Conn) { - ticker := time.NewTicker(HeartbeatInterval) - defer ticker.Stop() - for { - select { - case <-ctx.Done(): - return - case <-ticker.C: - } - err := conn.Ping(ctx) - if err != nil { - return - } - } -} - -// Heartbeat loops to ping a WebSocket to keep it alive. It calls `exit` on ping +// HeartbeatClose loops to ping a WebSocket to keep it alive. It calls `exit` on ping // failure. func HeartbeatClose(ctx context.Context, logger slog.Logger, exit func(), conn *websocket.Conn) { ticker := time.NewTicker(HeartbeatInterval) diff --git a/coderd/parameters.go b/coderd/parameters.go index cb24dcd431..08d571267a 100644 --- a/coderd/parameters.go +++ b/coderd/parameters.go @@ -139,7 +139,7 @@ func (api *API) handleParameterWebsocket(rw http.ResponseWriter, r *http.Request }) return } - go httpapi.Heartbeat(ctx, conn) + go httpapi.HeartbeatClose(ctx, api.Logger, cancel, conn) stream := wsjson.NewStream[codersdk.DynamicParametersRequest, codersdk.DynamicParametersResponse]( conn, diff --git a/coderd/provisionerjobs.go b/coderd/provisionerjobs.go index d93571644a..15b96d40df 100644 --- a/coderd/provisionerjobs.go +++ b/coderd/provisionerjobs.go @@ -544,7 +544,7 @@ func (f *logFollower) follow() { return } defer f.conn.Close(websocket.StatusNormalClosure, "done") - go httpapi.Heartbeat(f.ctx, f.conn) + go httpapi.HeartbeatClose(f.ctx, f.logger, cancel, f.conn) f.enc = wsjson.NewEncoder[codersdk.ProvisionerJobLog](f.conn, websocket.MessageText) // query for logs once right away, so we can get historical data from before diff --git a/coderd/workspaceagents.go b/coderd/workspaceagents.go index 68835c19c5..38e63be3d0 100644 --- a/coderd/workspaceagents.go +++ b/coderd/workspaceagents.go @@ -612,7 +612,9 @@ func (api *API) workspaceAgentLogs(rw http.ResponseWriter, r *http.Request) { }) return } - go httpapi.Heartbeat(ctx, conn) + ctx, cancel := context.WithCancel(ctx) + defer cancel() + go httpapi.HeartbeatClose(ctx, api.Logger, cancel, conn) encoder := wsjson.NewEncoder[[]codersdk.WorkspaceAgentLog](conn, websocket.MessageText) defer encoder.Close(websocket.StatusNormalClosure) @@ -1477,7 +1479,9 @@ func (api *API) workspaceAgentClientCoordinate(rw http.ResponseWriter, r *http.R ctx, wsNetConn := codersdk.WebsocketNetConn(ctx, conn, websocket.MessageBinary) defer wsNetConn.Close() - go httpapi.Heartbeat(ctx, conn) + ctx, cancel := context.WithCancel(ctx) + defer cancel() + go httpapi.HeartbeatClose(ctx, api.Logger, cancel, conn) defer conn.Close(websocket.StatusNormalClosure, "") err = api.TailnetClientService.ServeClient(ctx, version, wsNetConn, tailnet.StreamID{ @@ -1688,7 +1692,7 @@ func (api *API) watchWorkspaceAgentMetadataSSE(rw http.ResponseWriter, r *http.R // @Router /workspaceagents/{workspaceagent}/watch-metadata-ws [get] // @x-apidocgen {"skip": true} func (api *API) watchWorkspaceAgentMetadataWS(rw http.ResponseWriter, r *http.Request) { - api.watchWorkspaceAgentMetadata(rw, r, httpapi.OneWayWebSocketEventSender) + api.watchWorkspaceAgentMetadata(rw, r, httpapi.OneWayWebSocketEventSender(api.Logger)) } func (api *API) watchWorkspaceAgentMetadata( @@ -2287,7 +2291,9 @@ func (api *API) tailnetRPCConn(rw http.ResponseWriter, r *http.Request) { }) }() - go httpapi.Heartbeat(ctx, conn) + ctx, cancel := context.WithCancel(ctx) + defer cancel() + go httpapi.HeartbeatClose(ctx, api.Logger, cancel, conn) err = api.TailnetClientService.ServeClient(ctx, version, wsNetConn, tailnet.StreamID{ Name: "client", ID: peerID, diff --git a/coderd/workspaces.go b/coderd/workspaces.go index d368ce1f8f..ff179b24f5 100644 --- a/coderd/workspaces.go +++ b/coderd/workspaces.go @@ -2028,7 +2028,7 @@ func (api *API) watchWorkspaceSSE(rw http.ResponseWriter, r *http.Request) { // @Success 200 {object} codersdk.ServerSentEvent // @Router /workspaces/{workspace}/watch-ws [get] func (api *API) watchWorkspaceWS(rw http.ResponseWriter, r *http.Request) { - api.watchWorkspace(rw, r, httpapi.OneWayWebSocketEventSender) + api.watchWorkspace(rw, r, httpapi.OneWayWebSocketEventSender(api.Logger)) } func (api *API) watchWorkspace(