mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
chore: replace httpapi.Heartbeat with httpapi.HeartbeatClose (#21676)
Relates to https://github.com/coder/coder/pull/21676 * Replaces all existing usages of `httpapi.Heartbeat` with `httpapi.HeartbeatClose` * Removes `httpapi.HeartbeatClose`
This commit is contained in:
@@ -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)
|
||||
|
||||
|
||||
+71
-64
@@ -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.
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user