mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: add coderd_api_websocket_probes_total metric (#25012)
Relates to CODAGT-115 Adds metric `coderd_api_websocket_probes_total`. Every successful heartbeat for a given path will increment the metric. Comparing this with `coderd_api_concurrent_websockets` will give an indication of how many websocket connections are open but in a 'wedged' state (when heartbeats stopped versus when we closed the connection).
This commit is contained in:
+6
-1
@@ -914,6 +914,9 @@ func New(options *Options) *API {
|
||||
options.WorkspaceAppsStatsCollectorOptions.Reporter = api.statsReporter
|
||||
}
|
||||
|
||||
wsMetrics := httpmw.NewWSMetrics(options.PrometheusRegistry)
|
||||
api.wsWatcher = httpapi.NewWSWatcher(options.Clock, wsMetrics.RecordProbe)
|
||||
|
||||
api.workspaceAppServer = workspaceapps.NewServer(workspaceapps.ServerOptions{
|
||||
Logger: workspaceAppsLogger,
|
||||
|
||||
@@ -926,6 +929,7 @@ func New(options *Options) *API {
|
||||
SignedTokenProvider: api.WorkspaceAppsProvider,
|
||||
AgentProvider: api.agentProvider,
|
||||
StatsCollector: workspaceapps.NewStatsCollector(options.WorkspaceAppsStatsCollectorOptions),
|
||||
WSWatcher: api.wsWatcher,
|
||||
|
||||
DisablePathApps: options.DeploymentValues.DisablePathApps.Value(),
|
||||
CookiesConfig: options.DeploymentValues.HTTPCookies,
|
||||
@@ -994,7 +998,7 @@ func New(options *Options) *API {
|
||||
options.PrometheusRegistry.MustRegister(derpmetrics.NewDERPExpvarCollector(options.DERPServer))
|
||||
}
|
||||
cors := httpmw.Cors(options.DeploymentValues.Dangerous.AllowAllCors.Value())
|
||||
prometheusMW := httpmw.Prometheus(options.PrometheusRegistry)
|
||||
prometheusMW := httpmw.Prometheus(options.PrometheusRegistry, wsMetrics)
|
||||
|
||||
r.Use(
|
||||
sharedhttpmw.Recover(api.Logger),
|
||||
@@ -2251,6 +2255,7 @@ type API struct {
|
||||
metadataBatcher *metadatabatcher.Batcher
|
||||
lifecycleMetrics *agentapi.LifecycleMetrics
|
||||
workspaceAgentRPCMetrics *WorkspaceAgentRPCMetrics
|
||||
wsWatcher *httpapi.WSWatcher
|
||||
|
||||
Acquirer *provisionerdserver.Acquirer
|
||||
// dbRolluper rolls up template usage stats from raw agent and app
|
||||
|
||||
@@ -26,6 +26,7 @@ import (
|
||||
"github.com/coder/coder/v2/coderd/coderdtest"
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/dbfake"
|
||||
"github.com/coder/coder/v2/coderd/httpapi"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/codersdk/workspacesdk"
|
||||
"github.com/coder/coder/v2/provisioner/echo"
|
||||
@@ -33,6 +34,8 @@ import (
|
||||
"github.com/coder/coder/v2/tailnet"
|
||||
tailnetproto "github.com/coder/coder/v2/tailnet/proto"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
"github.com/coder/quartz"
|
||||
"github.com/coder/websocket"
|
||||
)
|
||||
|
||||
// updateGoldenFiles is a flag that can be set to update golden files.
|
||||
@@ -436,6 +439,69 @@ func TestDERPMetrics(t *testing.T) {
|
||||
"expected coder_derp_server_packets_dropped_reason_total to be registered")
|
||||
}
|
||||
|
||||
// TestWebSocketProbeMetrics verifies that the coderd_api_websocket_probes_total
|
||||
// metric is recorded end-to-end through a real coderd server.
|
||||
func TestWebSocketProbeMetrics(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitLong)
|
||||
mClock := quartz.NewMock(t)
|
||||
|
||||
trap := mClock.Trap().NewTicker("WSWatcher")
|
||||
defer trap.Close()
|
||||
|
||||
client, _, api := coderdtest.NewWithAPI(t, &coderdtest.Options{
|
||||
Clock: mClock,
|
||||
})
|
||||
firstUser := coderdtest.CreateFirstUser(t, client)
|
||||
member, _ := coderdtest.CreateAnotherUser(t, client, firstUser.OrganizationID)
|
||||
|
||||
// Open a WebSocket connection to the inbox watch endpoint.
|
||||
u, err := member.URL.Parse("/api/v2/notifications/inbox/watch")
|
||||
require.NoError(t, err)
|
||||
|
||||
// nolint:bodyclose
|
||||
wsConn, resp, err := websocket.Dial(ctx, u.String(), &websocket.DialOptions{
|
||||
HTTPHeader: http.Header{
|
||||
"Coder-Session-Token": []string{member.SessionToken()},
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
if resp != nil && resp.StatusCode != http.StatusSwitchingProtocols {
|
||||
err = codersdk.ReadBodyAsError(resp)
|
||||
}
|
||||
require.NoError(t, err)
|
||||
}
|
||||
defer wsConn.Close(websocket.StatusNormalClosure, "done")
|
||||
|
||||
// Start a reader to process control frames (pong responses).
|
||||
go func() {
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
default:
|
||||
_, _, err := wsConn.Read(ctx)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
// Wait for the WSWatcher ticker to be created, then trigger one probe.
|
||||
trap.MustWait(ctx).MustRelease(ctx)
|
||||
mClock.Advance(httpapi.HeartbeatInterval).MustWait(ctx)
|
||||
|
||||
// Assert the probe metric was recorded.
|
||||
testutil.Eventually(ctx, t, func(context.Context) bool {
|
||||
metrics, err := api.Options.PrometheusRegistry.Gather()
|
||||
assert.NoError(t, err)
|
||||
return testutil.PromCounterHasValue(t, metrics, 1,
|
||||
"coderd_api_websocket_probes_total", "/api/v2/notifications/inbox/watch", "ok")
|
||||
}, testutil.IntervalFast, "websocket probe metric not recorded")
|
||||
}
|
||||
|
||||
// TestRateLimitByUser verifies that rate limiting keys by user ID when
|
||||
// an authenticated session is present, rather than falling back to IP.
|
||||
// This is a regression test for https://github.com/coder/coder/issues/20857
|
||||
|
||||
+4
-5
@@ -192,7 +192,7 @@ func (api *API) watchChats(rw http.ResponseWriter, r *http.Request) {
|
||||
ctx, wsNetConn := codersdk.WebsocketNetConn(ctx, conn, websocket.MessageText)
|
||||
defer wsNetConn.Close()
|
||||
|
||||
go httpapi.HeartbeatClose(ctx, logger, cancel, conn)
|
||||
ctx = api.wsWatcher.Watch(ctx, logger, conn)
|
||||
|
||||
// The encoder is only written from the SubscribeWithErr callback,
|
||||
// which delivers serially per subscription. Do not add a second
|
||||
@@ -2393,8 +2393,7 @@ func (api *API) watchChatGit(rw http.ResponseWriter, r *http.Request) {
|
||||
|
||||
ctx, cancel := context.WithCancel(r.Context())
|
||||
defer cancel()
|
||||
|
||||
go httpapi.HeartbeatClose(ctx, logger, cancel, clientConn)
|
||||
ctx = api.wsWatcher.Watch(ctx, logger, clientConn)
|
||||
|
||||
// Proxy agent → client.
|
||||
agentCh := agentStream.Chan()
|
||||
@@ -2551,7 +2550,7 @@ func (api *API) watchChatDesktop(rw http.ResponseWriter, r *http.Request) {
|
||||
ctx, wsNetConn := workspaceapps.WebsocketNetConn(ctx, conn, websocket.MessageBinary)
|
||||
defer wsNetConn.Close()
|
||||
|
||||
go httpapi.HeartbeatClose(ctx, logger, cancel, conn)
|
||||
ctx = api.wsWatcher.Watch(ctx, logger, conn)
|
||||
|
||||
agentssh.Bicopy(ctx, wsNetConn, desktopConn)
|
||||
logger.Debug(ctx, "desktop Bicopy finished")
|
||||
@@ -3502,7 +3501,7 @@ func (api *API) streamChat(rw http.ResponseWriter, r *http.Request) {
|
||||
ctx, wsNetConn := codersdk.WebsocketNetConn(ctx, conn, websocket.MessageText)
|
||||
defer wsNetConn.Close()
|
||||
|
||||
go httpapi.HeartbeatClose(ctx, logger, cancel, conn)
|
||||
ctx = api.wsWatcher.Watch(ctx, logger, conn)
|
||||
|
||||
// The last_read_message_id field is owner-scoped. Shared readers
|
||||
// intentionally lack chat update permission, so their streams must not
|
||||
|
||||
@@ -419,7 +419,7 @@ 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(log slog.Logger) func(rw http.ResponseWriter, r *http.Request) (
|
||||
func OneWayWebSocketEventSender(log slog.Logger, watcher *WSWatcher) func(rw http.ResponseWriter, r *http.Request) (
|
||||
func(event codersdk.ServerSentEvent) error,
|
||||
<-chan struct{},
|
||||
error,
|
||||
@@ -436,7 +436,7 @@ func OneWayWebSocketEventSender(log slog.Logger) func(rw http.ResponseWriter, r
|
||||
cancel()
|
||||
return nil, nil, xerrors.Errorf("cannot establish connection: %w", err)
|
||||
}
|
||||
go HeartbeatClose(ctx, log, cancel, socket)
|
||||
ctx = watcher.Watch(ctx, log, socket)
|
||||
|
||||
eventC := make(chan codersdk.ServerSentEvent, 64)
|
||||
socketErrC := make(chan websocket.CloseError, 1)
|
||||
|
||||
@@ -22,6 +22,7 @@ import (
|
||||
"github.com/coder/coder/v2/coderd/httpapi"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
"github.com/coder/quartz"
|
||||
)
|
||||
|
||||
func TestInternalServerError(t *testing.T) {
|
||||
@@ -245,7 +246,7 @@ func TestOneWayWebSocketEventSender(t *testing.T) {
|
||||
req.Proto = p.proto
|
||||
|
||||
writer := newOneWayWriter(t)
|
||||
_, _, err := httpapi.OneWayWebSocketEventSender(slogtest.Make(t, nil))(writer, req)
|
||||
_, _, err := httpapi.OneWayWebSocketEventSender(slogtest.Make(t, nil), nil)(writer, req)
|
||||
require.ErrorContains(t, err, p.proto)
|
||||
}
|
||||
})
|
||||
@@ -254,9 +255,11 @@ func TestOneWayWebSocketEventSender(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
wsw := httpapi.NewWSWatcher(quartz.NewReal(), nil)
|
||||
|
||||
req := newBaseRequest(ctx)
|
||||
writer := newOneWayWriter(t)
|
||||
send, _, err := httpapi.OneWayWebSocketEventSender(slogtest.Make(t, nil))(writer, req)
|
||||
send, _, err := httpapi.OneWayWebSocketEventSender(slogtest.Make(t, nil), wsw)(writer, req)
|
||||
require.NoError(t, err)
|
||||
|
||||
serverPayload := codersdk.ServerSentEvent{
|
||||
@@ -280,9 +283,10 @@ func TestOneWayWebSocketEventSender(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx, cancel := context.WithCancel(testutil.Context(t, testutil.WaitShort))
|
||||
wsw := httpapi.NewWSWatcher(quartz.NewReal(), nil)
|
||||
req := newBaseRequest(ctx)
|
||||
writer := newOneWayWriter(t)
|
||||
_, done, err := httpapi.OneWayWebSocketEventSender(slogtest.Make(t, nil))(writer, req)
|
||||
_, done, err := httpapi.OneWayWebSocketEventSender(slogtest.Make(t, nil), wsw)(writer, req)
|
||||
require.NoError(t, err)
|
||||
|
||||
successC := make(chan bool)
|
||||
@@ -304,9 +308,10 @@ func TestOneWayWebSocketEventSender(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
wsw := httpapi.NewWSWatcher(quartz.NewReal(), nil)
|
||||
req := newBaseRequest(ctx)
|
||||
writer := newOneWayWriter(t)
|
||||
_, done, err := httpapi.OneWayWebSocketEventSender(slogtest.Make(t, nil))(writer, req)
|
||||
_, done, err := httpapi.OneWayWebSocketEventSender(slogtest.Make(t, nil), wsw)(writer, req)
|
||||
require.NoError(t, err)
|
||||
|
||||
successC := make(chan bool)
|
||||
@@ -334,9 +339,10 @@ func TestOneWayWebSocketEventSender(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx, cancel := context.WithCancel(testutil.Context(t, testutil.WaitShort))
|
||||
wsw := httpapi.NewWSWatcher(quartz.NewReal(), nil)
|
||||
req := newBaseRequest(ctx)
|
||||
writer := newOneWayWriter(t)
|
||||
send, done, err := httpapi.OneWayWebSocketEventSender(slogtest.Make(t, nil))(writer, req)
|
||||
send, done, err := httpapi.OneWayWebSocketEventSender(slogtest.Make(t, nil), wsw)(writer, req)
|
||||
require.NoError(t, err)
|
||||
|
||||
successC := make(chan bool)
|
||||
@@ -375,9 +381,10 @@ func TestOneWayWebSocketEventSender(t *testing.T) {
|
||||
timeout := hbDuration + (5 * time.Second)
|
||||
|
||||
ctx := testutil.Context(t, timeout)
|
||||
wsw := httpapi.NewWSWatcher(quartz.NewReal(), nil)
|
||||
req := newBaseRequest(ctx)
|
||||
writer := newOneWayWriter(t)
|
||||
_, _, err := httpapi.OneWayWebSocketEventSender(slogtest.Make(t, nil))(writer, req)
|
||||
_, _, err := httpapi.OneWayWebSocketEventSender(slogtest.Make(t, nil), wsw)(writer, req)
|
||||
require.NoError(t, err)
|
||||
|
||||
type Result struct {
|
||||
|
||||
+103
-39
@@ -15,20 +15,70 @@ import (
|
||||
|
||||
const HeartbeatInterval time.Duration = 15 * time.Second
|
||||
|
||||
// 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) {
|
||||
heartbeatCloseWith(ctx, logger, exit, conn, quartz.NewReal(), HeartbeatInterval)
|
||||
// ProbeResult classifies the outcome of a single WebSocket liveness
|
||||
// probe so that callers (typically a Prometheus recorder) can track
|
||||
// successes and the various failure modes independently.
|
||||
type ProbeResult string
|
||||
|
||||
const (
|
||||
ProbeOK ProbeResult = "ok"
|
||||
ProbeTimeout ProbeResult = "timeout"
|
||||
ProbePeerClosed ProbeResult = "peer_closed"
|
||||
ProbeCanceled ProbeResult = "canceled"
|
||||
ProbeError ProbeResult = "error"
|
||||
)
|
||||
|
||||
// ProbeRecorder is called once per liveness probe with its outcome.
|
||||
// It may be nil, in which case probes are still run but not recorded.
|
||||
type ProbeRecorder func(ctx context.Context, result ProbeResult)
|
||||
|
||||
// PingCloser is the minimal interface for WebSocket liveness probing.
|
||||
// *websocket.Conn satisfies this interface.
|
||||
type PingCloser interface {
|
||||
Ping(ctx context.Context) error
|
||||
Close(code websocket.StatusCode, reason string) error
|
||||
}
|
||||
|
||||
// HeartbeatCloseWithClock is like HeartbeatClose, but uses the provided
|
||||
// clock so tests can drive heartbeat ticks deterministically.
|
||||
func HeartbeatCloseWithClock(ctx context.Context, logger slog.Logger, exit func(), conn *websocket.Conn, clk quartz.Clock) {
|
||||
heartbeatCloseWith(ctx, logger, exit, conn, clk, HeartbeatInterval)
|
||||
// WSWatcher supervises WebSocket connections for liveness by
|
||||
// periodically sending ping frames. On probe failure, the watcher
|
||||
// closes the connection with StatusGoingAway and cancels the
|
||||
// returned context; the caller owns closing the connection on
|
||||
// normal teardown.
|
||||
type WSWatcher struct {
|
||||
rec ProbeRecorder
|
||||
clk quartz.Clock
|
||||
interval time.Duration
|
||||
}
|
||||
|
||||
func heartbeatCloseWith(ctx context.Context, logger slog.Logger, exit func(), conn *websocket.Conn, clk quartz.Clock, interval time.Duration) {
|
||||
ticker := clk.NewTicker(interval, "HeartbeatClose")
|
||||
// NewWSWatcher creates a WSWatcher. Pass nil for rec when no
|
||||
// recording is needed (e.g. agent-side code without a Prometheus
|
||||
// registry).
|
||||
func NewWSWatcher(clk quartz.Clock, rec ProbeRecorder) *WSWatcher {
|
||||
return &WSWatcher{
|
||||
rec: rec,
|
||||
clk: clk,
|
||||
interval: HeartbeatInterval,
|
||||
}
|
||||
}
|
||||
|
||||
// Watch supervises conn for liveness. The returned context is
|
||||
// canceled when parent is canceled or when conn fails a probe.
|
||||
// Watch closes conn on probe failure with StatusGoingAway; the
|
||||
// caller owns close on normal teardown.
|
||||
func (w *WSWatcher) Watch(parent context.Context, log slog.Logger, conn PingCloser) context.Context {
|
||||
if w == nil {
|
||||
panic("developer error: WSWatcher is nil")
|
||||
}
|
||||
ctx, cancel := context.WithCancel(parent)
|
||||
go func() {
|
||||
defer cancel()
|
||||
w.supervise(ctx, log, conn)
|
||||
}()
|
||||
return ctx
|
||||
}
|
||||
|
||||
func (w *WSWatcher) supervise(ctx context.Context, log slog.Logger, conn PingCloser) {
|
||||
ticker := w.clk.NewTicker(w.interval, "WSWatcher")
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
@@ -37,39 +87,53 @@ func heartbeatCloseWith(ctx context.Context, logger slog.Logger, exit func(), co
|
||||
return
|
||||
case <-ticker.C:
|
||||
}
|
||||
err := pingWithTimeout(ctx, conn, interval)
|
||||
if err != nil {
|
||||
// These errors are all expected during normal connection
|
||||
// teardown and should not be logged at error level:
|
||||
// - context.DeadlineExceeded: client disconnected
|
||||
// without sending a close frame.
|
||||
// - context.Canceled: request context was canceled.
|
||||
// - net.ErrClosed: connection was already closed by
|
||||
// another goroutine (e.g. handler returned).
|
||||
// - websocket.CloseError: a close frame was
|
||||
// received or sent.
|
||||
if errors.Is(err, context.DeadlineExceeded) ||
|
||||
errors.Is(err, context.Canceled) ||
|
||||
errors.Is(err, net.ErrClosed) ||
|
||||
websocket.CloseStatus(err) != -1 {
|
||||
logger.Debug(ctx, "heartbeat ping stopped", slog.Error(err))
|
||||
} else {
|
||||
logger.Error(ctx, "failed to heartbeat ping", slog.Error(err))
|
||||
}
|
||||
_ = conn.Close(websocket.StatusGoingAway, "Ping failed")
|
||||
exit()
|
||||
return
|
||||
|
||||
result, err := probe(ctx, conn, w.interval)
|
||||
if w.rec != nil {
|
||||
w.rec(ctx, result)
|
||||
}
|
||||
if result == ProbeOK {
|
||||
continue
|
||||
}
|
||||
if result == ProbeError {
|
||||
log.Error(ctx, "websocket probe failed", slog.Error(err))
|
||||
} else {
|
||||
log.Debug(ctx, "websocket probe stopped",
|
||||
slog.F("result", string(result)), slog.Error(err))
|
||||
}
|
||||
_ = conn.Close(websocket.StatusGoingAway, "liveness probe failed")
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
func pingWithTimeout(ctx context.Context, conn *websocket.Conn, timeout time.Duration) error {
|
||||
ctx, cancel := context.WithTimeout(ctx, timeout)
|
||||
func probe(ctx context.Context, conn PingCloser, timeout time.Duration) (ProbeResult, error) {
|
||||
pingCtx, cancel := context.WithTimeout(ctx, timeout)
|
||||
defer cancel()
|
||||
err := conn.Ping(ctx)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("failed to ping: %w", err)
|
||||
err := conn.Ping(pingCtx)
|
||||
switch {
|
||||
case err == nil:
|
||||
return ProbeOK, nil
|
||||
case errors.Is(err, context.Canceled):
|
||||
return ProbeCanceled, err
|
||||
case errors.Is(err, context.DeadlineExceeded):
|
||||
return ProbeTimeout, err
|
||||
case errors.Is(err, net.ErrClosed) || websocket.CloseStatus(err) != -1:
|
||||
return ProbePeerClosed, err
|
||||
default:
|
||||
return ProbeError, xerrors.Errorf("ping: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
// HeartbeatClose is a legacy helper that pings conn in a loop and
|
||||
// calls exit on failure. Callers that need metric recording should
|
||||
// use WSWatcher directly.
|
||||
func HeartbeatClose(ctx context.Context, logger slog.Logger, exit func(), conn *websocket.Conn) {
|
||||
w := NewWSWatcher(quartz.NewReal(), nil)
|
||||
watchCtx := w.Watch(ctx, logger, conn)
|
||||
<-watchCtx.Done()
|
||||
// Only call exit when the probe failed; if the parent context was
|
||||
// canceled the caller is already shutting down.
|
||||
if ctx.Err() == nil {
|
||||
exit()
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -4,11 +4,14 @@ import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/prometheus/client_golang/prometheus"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"cdr.dev/slog/v3"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
@@ -53,7 +56,37 @@ func websocketPair(ctx context.Context, t *testing.T) *websocket.Conn {
|
||||
}
|
||||
}
|
||||
|
||||
func TestHeartbeatClose(t *testing.T) {
|
||||
// probeRecords is a thread-safe collector for ProbeResult values.
|
||||
type probeRecords struct {
|
||||
mu sync.Mutex
|
||||
results []ProbeResult
|
||||
}
|
||||
|
||||
func (r *probeRecords) record(_ context.Context, result ProbeResult) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.results = append(r.results, result)
|
||||
}
|
||||
|
||||
func (r *probeRecords) count(want ProbeResult) int {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
n := 0
|
||||
for _, got := range r.results {
|
||||
if got == want {
|
||||
n++
|
||||
}
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
func (r *probeRecords) len() int {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
return len(r.results)
|
||||
}
|
||||
|
||||
func TestWSWatcher(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("ServerSideClose", func(t *testing.T) {
|
||||
@@ -63,33 +96,31 @@ func TestHeartbeatClose(t *testing.T) {
|
||||
sink := testutil.NewFakeSink(t)
|
||||
logger := sink.Logger()
|
||||
mClock := quartz.NewMock(t)
|
||||
rec := &probeRecords{}
|
||||
|
||||
// Trap ticker creation so we can synchronize startup.
|
||||
trap := mClock.Trap().NewTicker("HeartbeatClose")
|
||||
trap := mClock.Trap().NewTicker("WSWatcher")
|
||||
defer trap.Close()
|
||||
|
||||
serverConn := websocketPair(ctx, t)
|
||||
exitCalled := make(chan struct{})
|
||||
|
||||
go heartbeatCloseWith(ctx, logger, func() {
|
||||
close(exitCalled)
|
||||
}, serverConn, mClock, time.Second)
|
||||
w := &WSWatcher{rec: rec.record, clk: mClock, interval: time.Second}
|
||||
watchCtx := w.Watch(ctx, logger, serverConn)
|
||||
|
||||
// Wait for the ticker to be created, then release.
|
||||
trap.MustWait(ctx).MustRelease(ctx)
|
||||
|
||||
// Close the server-side connection before the tick fires.
|
||||
// The next ping will get net.ErrClosed.
|
||||
// The next ping will get a close/net.ErrClosed error.
|
||||
_ = serverConn.Close(websocket.StatusGoingAway, "simulated teardown")
|
||||
|
||||
// Advance clock to trigger the tick.
|
||||
mClock.Advance(time.Second).MustWait(ctx)
|
||||
|
||||
// Wait for heartbeatClose to call exit.
|
||||
// The watch context should be canceled after probe failure.
|
||||
select {
|
||||
case <-exitCalled:
|
||||
case <-watchCtx.Done():
|
||||
case <-ctx.Done():
|
||||
t.Fatal("timed out waiting for heartbeatClose to call exit")
|
||||
t.Fatal("timed out waiting for watch context to be canceled")
|
||||
}
|
||||
|
||||
// A closed connection is a normal shutdown condition. The
|
||||
@@ -100,6 +131,9 @@ func TestHeartbeatClose(t *testing.T) {
|
||||
debugEntries := sink.Entries(func(e slog.SinkEntry) bool { return e.Level == slog.LevelDebug })
|
||||
assert.NotEmpty(t, debugEntries,
|
||||
"expected a debug-level log entry for the closed connection")
|
||||
assert.Zero(t, rec.count(ProbeOK), "expected no successful probes")
|
||||
assert.Equal(t, 1, rec.len(), "expected exactly one probe recorded")
|
||||
assert.Equal(t, 1, rec.count(ProbePeerClosed), "expected one peer_closed probe")
|
||||
})
|
||||
|
||||
t.Run("ContextCanceled", func(t *testing.T) {
|
||||
@@ -109,36 +143,33 @@ func TestHeartbeatClose(t *testing.T) {
|
||||
sink := testutil.NewFakeSink(t)
|
||||
logger := sink.Logger()
|
||||
mClock := quartz.NewMock(t)
|
||||
rec := &probeRecords{}
|
||||
|
||||
trap := mClock.Trap().NewTicker("HeartbeatClose")
|
||||
trap := mClock.Trap().NewTicker("WSWatcher")
|
||||
defer trap.Close()
|
||||
|
||||
serverCtx, serverCancel := context.WithCancel(ctx)
|
||||
serverConn := websocketPair(ctx, t)
|
||||
done := make(chan struct{})
|
||||
|
||||
go func() {
|
||||
defer close(done)
|
||||
heartbeatCloseWith(serverCtx, logger, func() {
|
||||
t.Error("exit should not be called on context cancel")
|
||||
}, serverConn, mClock, time.Second)
|
||||
}()
|
||||
w := &WSWatcher{rec: rec.record, clk: mClock, interval: time.Second}
|
||||
watchCtx := w.Watch(serverCtx, logger, serverConn)
|
||||
|
||||
trap.MustWait(ctx).MustRelease(ctx)
|
||||
|
||||
// Cancel the context. HeartbeatClose should return via
|
||||
// the <-ctx.Done() branch without calling exit.
|
||||
// Cancel the parent context. The watcher should exit via
|
||||
// the <-ctx.Done() branch without closing the conn.
|
||||
serverCancel()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
case <-watchCtx.Done():
|
||||
case <-ctx.Done():
|
||||
t.Fatal("timed out waiting for heartbeatClose to return")
|
||||
t.Fatal("timed out waiting for watch context to be canceled")
|
||||
}
|
||||
|
||||
errorEntries := sink.Entries(func(e slog.SinkEntry) bool { return e.Level == slog.LevelError })
|
||||
assert.Empty(t, errorEntries,
|
||||
"context cancellation should not produce error-level logs, got: %+v", errorEntries)
|
||||
assert.Zero(t, rec.len(), "expected no probes when context is canceled before tick")
|
||||
})
|
||||
|
||||
t.Run("PingSucceeds", func(t *testing.T) {
|
||||
@@ -148,30 +179,30 @@ func TestHeartbeatClose(t *testing.T) {
|
||||
sink := testutil.NewFakeSink(t)
|
||||
logger := sink.Logger()
|
||||
mClock := quartz.NewMock(t)
|
||||
rec := &probeRecords{}
|
||||
|
||||
trap := mClock.Trap().NewTicker("HeartbeatClose")
|
||||
trap := mClock.Trap().NewTicker("WSWatcher")
|
||||
defer trap.Close()
|
||||
|
||||
serverConn := websocketPair(ctx, t)
|
||||
exitCalled := make(chan struct{}, 1)
|
||||
|
||||
go heartbeatCloseWith(ctx, logger, func() {
|
||||
exitCalled <- struct{}{}
|
||||
}, serverConn, mClock, time.Second)
|
||||
w := &WSWatcher{rec: rec.record, clk: mClock, interval: time.Second}
|
||||
watchCtx := w.Watch(ctx, logger, serverConn)
|
||||
|
||||
trap.MustWait(ctx).MustRelease(ctx)
|
||||
|
||||
// Fire several ticks — pings should succeed each time.
|
||||
for range 3 {
|
||||
// Fire several ticks; pings should succeed each time.
|
||||
for i := range 3 {
|
||||
mClock.Advance(time.Second).MustWait(ctx)
|
||||
|
||||
// Give the ping round-trip time to complete.
|
||||
// If exit were called, we'd catch it.
|
||||
select {
|
||||
case <-exitCalled:
|
||||
t.Fatal("exit should not be called when pings succeed")
|
||||
default:
|
||||
}
|
||||
testutil.Eventually(ctx, t, func(context.Context) bool {
|
||||
select {
|
||||
case <-watchCtx.Done():
|
||||
t.Fatal("watch context should not be canceled when pings succeed")
|
||||
default:
|
||||
}
|
||||
return rec.count(ProbeOK) == i+1
|
||||
}, testutil.IntervalFast, "probe counter not incremented at tick %d", i+1)
|
||||
}
|
||||
|
||||
// No logs should be emitted during normal operation.
|
||||
@@ -181,5 +212,183 @@ func TestHeartbeatClose(t *testing.T) {
|
||||
debugEntries := sink.Entries(func(e slog.SinkEntry) bool { return e.Level == slog.LevelDebug })
|
||||
assert.Empty(t, debugEntries,
|
||||
"successful pings should not produce debug-level logs, got: %+v", debugEntries)
|
||||
assert.Equal(t, 3, rec.count(ProbeOK), "expected 3 successful probes")
|
||||
})
|
||||
|
||||
t.Run("RecordsPrometheusCounter", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
|
||||
// Use a real prometheus registry to verify end-to-end metric recording.
|
||||
registry := prometheus.NewRegistry()
|
||||
probes := prometheus.NewCounterVec(prometheus.CounterOpts{
|
||||
Namespace: "coderd",
|
||||
Subsystem: "api",
|
||||
Name: "websocket_probes_total",
|
||||
Help: "test",
|
||||
}, []string{"path", "result"})
|
||||
registry.MustRegister(probes)
|
||||
|
||||
recorder := func(ctx context.Context, r ProbeResult) {
|
||||
probes.WithLabelValues("/test/path", string(r)).Inc()
|
||||
}
|
||||
|
||||
sink := testutil.NewFakeSink(t)
|
||||
logger := sink.Logger()
|
||||
mClock := quartz.NewMock(t)
|
||||
|
||||
trap := mClock.Trap().NewTicker("WSWatcher")
|
||||
defer trap.Close()
|
||||
|
||||
serverConn := websocketPair(ctx, t)
|
||||
|
||||
w := &WSWatcher{rec: recorder, clk: mClock, interval: time.Second}
|
||||
watchCtx := w.Watch(ctx, logger, serverConn)
|
||||
|
||||
trap.MustWait(ctx).MustRelease(ctx)
|
||||
mClock.Advance(time.Second).MustWait(ctx)
|
||||
|
||||
testutil.Eventually(ctx, t, func(context.Context) bool {
|
||||
select {
|
||||
case <-watchCtx.Done():
|
||||
t.Fatal("watch context should not be canceled when pings succeed")
|
||||
default:
|
||||
}
|
||||
metrics, err := registry.Gather()
|
||||
require.NoError(t, err)
|
||||
return testutil.PromCounterHasValue(t, metrics, 1,
|
||||
"coderd_api_websocket_probes_total", "/test/path", "ok")
|
||||
}, testutil.IntervalFast, "probe counter not incremented")
|
||||
})
|
||||
|
||||
t.Run("ProbeTimeout", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
|
||||
sink := testutil.NewFakeSink(t)
|
||||
logger := sink.Logger()
|
||||
mClock := quartz.NewMock(t)
|
||||
rec := &probeRecords{}
|
||||
|
||||
trap := mClock.Trap().NewTicker("WSWatcher")
|
||||
defer trap.Close()
|
||||
|
||||
// Set up a websocket pair manually. Do NOT call CloseRead
|
||||
// on the client so pong frames are never sent back.
|
||||
serverConnCh := make(chan *websocket.Conn, 1)
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
conn, err := websocket.Accept(w, r, nil)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
serverConnCh <- conn
|
||||
<-ctx.Done()
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
//nolint:bodyclose
|
||||
clientConn, _, err := websocket.Dial(ctx, srv.URL, nil)
|
||||
require.NoError(t, err)
|
||||
// Intentionally NOT calling clientConn.CloseRead, so pongs won't be processed.
|
||||
t.Cleanup(func() {
|
||||
_ = clientConn.Close(websocket.StatusNormalClosure, "test cleanup")
|
||||
})
|
||||
|
||||
var serverConn *websocket.Conn
|
||||
select {
|
||||
case sc := <-serverConnCh:
|
||||
_ = sc.CloseRead(ctx)
|
||||
serverConn = sc
|
||||
case <-ctx.Done():
|
||||
t.Fatal("timed out waiting for server websocket accept")
|
||||
}
|
||||
|
||||
// Use a very short interval so the real context.WithTimeout
|
||||
// inside probe() expires quickly when pongs aren't coming.
|
||||
w := &WSWatcher{rec: rec.record, clk: mClock, interval: time.Millisecond}
|
||||
watchCtx := w.Watch(ctx, logger, serverConn)
|
||||
|
||||
trap.MustWait(ctx).MustRelease(ctx)
|
||||
mClock.Advance(time.Millisecond).MustWait(ctx)
|
||||
|
||||
// Wait for the watch context to be canceled (probe failure).
|
||||
select {
|
||||
case <-watchCtx.Done():
|
||||
case <-ctx.Done():
|
||||
t.Fatal("timed out waiting for watch context to be canceled")
|
||||
}
|
||||
|
||||
assert.Equal(t, 1, rec.count(ProbeTimeout), "expected one timeout probe")
|
||||
// Timeout is an expected condition, should be Debug not Error.
|
||||
errorEntries := sink.Entries(func(e slog.SinkEntry) bool { return e.Level == slog.LevelError })
|
||||
assert.Empty(t, errorEntries,
|
||||
"probe timeout should not produce error-level logs, got: %+v", errorEntries)
|
||||
})
|
||||
|
||||
t.Run("ProbeError", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
|
||||
sink := testutil.NewFakeSink(t)
|
||||
logger := sink.Logger()
|
||||
mClock := quartz.NewMock(t)
|
||||
rec := &probeRecords{}
|
||||
|
||||
trap := mClock.Trap().NewTicker("WSWatcher")
|
||||
defer trap.Close()
|
||||
|
||||
fConn := &fakePingCloser{
|
||||
pingErr: xerrors.New("unexpected internal error"),
|
||||
}
|
||||
|
||||
w := &WSWatcher{rec: rec.record, clk: mClock, interval: time.Second}
|
||||
watchCtx := w.Watch(ctx, logger, fConn)
|
||||
|
||||
trap.MustWait(ctx).MustRelease(ctx)
|
||||
mClock.Advance(time.Second).MustWait(ctx)
|
||||
|
||||
// Wait for the watch context to be canceled (probe failure).
|
||||
select {
|
||||
case <-watchCtx.Done():
|
||||
case <-ctx.Done():
|
||||
t.Fatal("timed out waiting for watch context to be canceled")
|
||||
}
|
||||
|
||||
assert.Equal(t, 1, rec.count(ProbeError), "expected one error probe")
|
||||
// ProbeError should log at Error level (unlike other failures).
|
||||
errorEntries := sink.Entries(func(e slog.SinkEntry) bool {
|
||||
return e.Level == slog.LevelError
|
||||
})
|
||||
assert.NotEmpty(t, errorEntries, "ProbeError should produce error-level log")
|
||||
|
||||
// Connection should be closed with StatusGoingAway.
|
||||
fConn.mu.Lock()
|
||||
assert.True(t, fConn.closed, "connection should be closed on probe error")
|
||||
assert.Equal(t, websocket.StatusGoingAway, fConn.code)
|
||||
fConn.mu.Unlock()
|
||||
})
|
||||
}
|
||||
|
||||
// fakePingCloser is a test double for the pingCloser interface.
|
||||
type fakePingCloser struct {
|
||||
mu sync.Mutex
|
||||
pingErr error
|
||||
closed bool
|
||||
code websocket.StatusCode
|
||||
reason string
|
||||
}
|
||||
|
||||
func (f *fakePingCloser) Ping(context.Context) error {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
return f.pingErr
|
||||
}
|
||||
|
||||
func (f *fakePingCloser) Close(code websocket.StatusCode, reason string) error {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
f.closed = true
|
||||
f.code = code
|
||||
f.reason = reason
|
||||
return nil
|
||||
}
|
||||
|
||||
+61
-24
@@ -1,6 +1,7 @@
|
||||
package httpmw
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"time"
|
||||
@@ -12,7 +13,63 @@ import (
|
||||
"github.com/coder/coder/v2/coderd/tracing"
|
||||
)
|
||||
|
||||
func Prometheus(register prometheus.Registerer) func(http.Handler) http.Handler {
|
||||
// WSMetrics groups all WebSocket-related Prometheus metrics so they
|
||||
// can be created once and shared between the HTTP middleware and the
|
||||
// WSWatcher probe recorder.
|
||||
type WSMetrics struct {
|
||||
Concurrent *prometheus.GaugeVec
|
||||
Durations *prometheus.HistogramVec
|
||||
Probes *prometheus.CounterVec
|
||||
}
|
||||
|
||||
// NewWSMetrics registers and returns WebSocket metrics. The returned
|
||||
// struct is safe to pass to both Prometheus() and
|
||||
// WSMetrics.RecordProbe.
|
||||
func NewWSMetrics(reg prometheus.Registerer) *WSMetrics {
|
||||
factory := promauto.With(reg)
|
||||
return &WSMetrics{
|
||||
Concurrent: factory.NewGaugeVec(prometheus.GaugeOpts{
|
||||
Namespace: "coderd",
|
||||
Subsystem: "api",
|
||||
Name: "concurrent_websockets",
|
||||
Help: "The total number of concurrent API websockets.",
|
||||
}, []string{"path"}),
|
||||
Durations: factory.NewHistogramVec(prometheus.HistogramOpts{
|
||||
Namespace: "coderd",
|
||||
Subsystem: "api",
|
||||
Name: "websocket_durations_seconds",
|
||||
Help: "Websocket duration distribution of requests in seconds.",
|
||||
Buckets: []float64{
|
||||
0.001, // 1ms
|
||||
1,
|
||||
60, // 1 minute
|
||||
60 * 60, // 1 hour
|
||||
60 * 60 * 15, // 15 hours
|
||||
60 * 60 * 30, // 30 hours
|
||||
},
|
||||
}, []string{"path"}),
|
||||
Probes: factory.NewCounterVec(prometheus.CounterOpts{
|
||||
Namespace: "coderd",
|
||||
Subsystem: "api",
|
||||
Name: "websocket_probes_total",
|
||||
Help: "WebSocket liveness probe outcomes by route. " +
|
||||
"Compare rate(...{result=\"ok\"}[1m]) against " +
|
||||
"coderd_api_concurrent_websockets to detect " +
|
||||
"unresponsive WebSocket connections.",
|
||||
}, []string{"path", "result"}),
|
||||
}
|
||||
}
|
||||
|
||||
// RecordProbe records a single liveness probe outcome. It extracts
|
||||
// the HTTP route from ctx via ExtractHTTPRoute.
|
||||
func (m *WSMetrics) RecordProbe(ctx context.Context, r httpapi.ProbeResult) {
|
||||
m.Probes.WithLabelValues(ExtractHTTPRoute(ctx), string(r)).Inc()
|
||||
}
|
||||
|
||||
func Prometheus(register prometheus.Registerer, ws *WSMetrics) func(http.Handler) http.Handler {
|
||||
if ws == nil {
|
||||
panic("developer error: WSMetrics is nil")
|
||||
}
|
||||
factory := promauto.With(register)
|
||||
requestsProcessed := factory.NewCounterVec(prometheus.CounterOpts{
|
||||
Namespace: "coderd",
|
||||
@@ -26,26 +83,6 @@ func Prometheus(register prometheus.Registerer) func(http.Handler) http.Handler
|
||||
Name: "concurrent_requests",
|
||||
Help: "The number of concurrent API requests.",
|
||||
}, []string{"method", "path"})
|
||||
websocketsConcurrent := factory.NewGaugeVec(prometheus.GaugeOpts{
|
||||
Namespace: "coderd",
|
||||
Subsystem: "api",
|
||||
Name: "concurrent_websockets",
|
||||
Help: "The total number of concurrent API websockets.",
|
||||
}, []string{"path"})
|
||||
websocketsDist := factory.NewHistogramVec(prometheus.HistogramOpts{
|
||||
Namespace: "coderd",
|
||||
Subsystem: "api",
|
||||
Name: "websocket_durations_seconds",
|
||||
Help: "Websocket duration distribution of requests in seconds.",
|
||||
Buckets: []float64{
|
||||
0.001, // 1ms
|
||||
1,
|
||||
60, // 1 minute
|
||||
60 * 60, // 1 hour
|
||||
60 * 60 * 15, // 15 hours
|
||||
60 * 60 * 30, // 30 hours
|
||||
},
|
||||
}, []string{"path"})
|
||||
requestsDist := factory.NewHistogramVec(prometheus.HistogramOpts{
|
||||
Namespace: "coderd",
|
||||
Subsystem: "api",
|
||||
@@ -74,10 +111,10 @@ func Prometheus(register prometheus.Registerer) func(http.Handler) http.Handler
|
||||
|
||||
// We want to count WebSockets separately.
|
||||
if httpapi.IsWebsocketUpgrade(r) {
|
||||
websocketsConcurrent.WithLabelValues(path).Inc()
|
||||
defer websocketsConcurrent.WithLabelValues(path).Dec()
|
||||
ws.Concurrent.WithLabelValues(path).Inc()
|
||||
defer ws.Concurrent.WithLabelValues(path).Dec()
|
||||
|
||||
dist = websocketsDist
|
||||
dist = ws.Durations
|
||||
} else {
|
||||
requestsConcurrent.WithLabelValues(method, path).Inc()
|
||||
defer requestsConcurrent.WithLabelValues(method, path).Dec()
|
||||
|
||||
@@ -29,7 +29,7 @@ func TestPrometheus(t *testing.T) {
|
||||
req = req.WithContext(context.WithValue(req.Context(), chi.RouteCtxKey, chi.NewRouteContext()))
|
||||
res := &tracing.StatusWriter{ResponseWriter: httptest.NewRecorder()}
|
||||
reg := prometheus.NewRegistry()
|
||||
httpmw.HTTPRoute(httpmw.Prometheus(reg)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
httpmw.HTTPRoute(httpmw.Prometheus(reg, httpmw.NewWSMetrics(reg))(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))).ServeHTTP(res, req)
|
||||
metrics, err := reg.Gather()
|
||||
@@ -43,7 +43,7 @@ func TestPrometheus(t *testing.T) {
|
||||
defer cancel()
|
||||
|
||||
reg := prometheus.NewRegistry()
|
||||
promMW := httpmw.Prometheus(reg)
|
||||
promMW := httpmw.Prometheus(reg, httpmw.NewWSMetrics(reg))
|
||||
|
||||
// Create a test handler to simulate a WebSocket connection
|
||||
testHandler := http.HandlerFunc(func(rw http.ResponseWriter, r *http.Request) {
|
||||
@@ -82,7 +82,7 @@ func TestPrometheus(t *testing.T) {
|
||||
t.Run("UserRoute", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
reg := prometheus.NewRegistry()
|
||||
promMW := httpmw.Prometheus(reg)
|
||||
promMW := httpmw.Prometheus(reg, httpmw.NewWSMetrics(reg))
|
||||
|
||||
r := chi.NewRouter()
|
||||
r.With(httpmw.HTTPRoute).With(promMW).Get("/api/v2/users/{user}", func(w http.ResponseWriter, r *http.Request) {})
|
||||
@@ -112,7 +112,7 @@ func TestPrometheus(t *testing.T) {
|
||||
t.Run("StaticRoute", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
reg := prometheus.NewRegistry()
|
||||
promMW := httpmw.Prometheus(reg)
|
||||
promMW := httpmw.Prometheus(reg, httpmw.NewWSMetrics(reg))
|
||||
|
||||
r := chi.NewRouter()
|
||||
r.Use(httpmw.HTTPRoute)
|
||||
@@ -143,7 +143,7 @@ func TestPrometheus(t *testing.T) {
|
||||
t.Run("UnknownRoute", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
reg := prometheus.NewRegistry()
|
||||
promMW := httpmw.Prometheus(reg)
|
||||
promMW := httpmw.Prometheus(reg, httpmw.NewWSMetrics(reg))
|
||||
|
||||
r := chi.NewRouter()
|
||||
r.Use(httpmw.HTTPRoute)
|
||||
@@ -172,7 +172,7 @@ func TestPrometheus(t *testing.T) {
|
||||
t.Run("Subrouter", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
reg := prometheus.NewRegistry()
|
||||
promMW := httpmw.Prometheus(reg)
|
||||
promMW := httpmw.Prometheus(reg, httpmw.NewWSMetrics(reg))
|
||||
|
||||
r := chi.NewRouter()
|
||||
r.Use(httpmw.HTTPRoute)
|
||||
|
||||
@@ -224,7 +224,7 @@ func (api *API) watchInboxNotifications(rw http.ResponseWriter, r *http.Request)
|
||||
ctx, wsNetConn := codersdk.WebsocketNetConn(ctx, conn, websocket.MessageText)
|
||||
defer wsNetConn.Close()
|
||||
|
||||
go httpapi.HeartbeatClose(ctx, logger, cancel, conn)
|
||||
ctx = api.wsWatcher.Watch(ctx, logger, conn)
|
||||
|
||||
encoder := json.NewEncoder(wsNetConn)
|
||||
|
||||
|
||||
@@ -140,7 +140,7 @@ func (api *API) handleParameterWebsocket(rw http.ResponseWriter, r *http.Request
|
||||
})
|
||||
return
|
||||
}
|
||||
go httpapi.HeartbeatClose(ctx, api.Logger, cancel, conn)
|
||||
ctx = api.wsWatcher.Watch(ctx, api.Logger, conn)
|
||||
|
||||
stream := wsjson.NewStream[codersdk.DynamicParametersRequest, codersdk.DynamicParametersResponse](
|
||||
conn,
|
||||
|
||||
+36
-29
@@ -202,7 +202,7 @@ func (api *API) provisionerJobLogs(rw http.ResponseWriter, r *http.Request, job
|
||||
return
|
||||
}
|
||||
|
||||
follower := newLogFollower(ctx, logger, api.Database, api.Pubsub, rw, r, job, after)
|
||||
follower := newLogFollower(ctx, logger, api.Database, api.Pubsub, api.wsWatcher, rw, r, job, after)
|
||||
api.WebsocketWaitMutex.Lock()
|
||||
api.WebsocketWaitGroup.Add(1)
|
||||
api.WebsocketWaitMutex.Unlock()
|
||||
@@ -493,14 +493,15 @@ func jobIsComplete(logger slog.Logger, job database.ProvisionerJob) bool {
|
||||
}
|
||||
|
||||
type logFollower struct {
|
||||
ctx context.Context
|
||||
logger slog.Logger
|
||||
db database.Store
|
||||
pubsub pubsub.Pubsub
|
||||
r *http.Request
|
||||
rw http.ResponseWriter
|
||||
conn *websocket.Conn
|
||||
enc *wsjson.Encoder[codersdk.ProvisionerJobLog]
|
||||
ctx context.Context
|
||||
logger slog.Logger
|
||||
db database.Store
|
||||
pubsub pubsub.Pubsub
|
||||
wsWatcher *httpapi.WSWatcher
|
||||
r *http.Request
|
||||
rw http.ResponseWriter
|
||||
conn *websocket.Conn
|
||||
enc *wsjson.Encoder[codersdk.ProvisionerJobLog]
|
||||
|
||||
jobID uuid.UUID
|
||||
after int64
|
||||
@@ -511,13 +512,15 @@ type logFollower struct {
|
||||
|
||||
func newLogFollower(
|
||||
ctx context.Context, logger slog.Logger, db database.Store, ps pubsub.Pubsub,
|
||||
rw http.ResponseWriter, r *http.Request, job database.ProvisionerJob, after int64,
|
||||
wsWatcher *httpapi.WSWatcher, rw http.ResponseWriter, r *http.Request,
|
||||
job database.ProvisionerJob, after int64,
|
||||
) *logFollower {
|
||||
return &logFollower{
|
||||
ctx: ctx,
|
||||
logger: logger,
|
||||
db: db,
|
||||
pubsub: ps,
|
||||
wsWatcher: wsWatcher,
|
||||
r: r,
|
||||
rw: rw,
|
||||
jobID: job.ID,
|
||||
@@ -579,26 +582,30 @@ func (f *logFollower) follow() {
|
||||
return
|
||||
}
|
||||
defer f.conn.Close(websocket.StatusNormalClosure, "done")
|
||||
go httpapi.HeartbeatClose(f.ctx, f.logger, cancel, f.conn)
|
||||
// Do not reassign f.ctx here; the listener method reads
|
||||
// f.ctx on the pubsub goroutine concurrently. Use a local
|
||||
// variable instead. The watched context is a child of f.ctx,
|
||||
// so canceling f.ctx still cascades.
|
||||
watchCtx := f.wsWatcher.Watch(f.ctx, f.logger, 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
|
||||
// subscription
|
||||
if err := f.query(); err != nil {
|
||||
if f.ctx.Err() == nil && !xerrors.Is(err, io.EOF) {
|
||||
if err := f.query(watchCtx); err != nil {
|
||||
if watchCtx.Err() == nil && !xerrors.Is(err, io.EOF) {
|
||||
// neither context expiry, nor EOF, close and log
|
||||
f.logger.Error(f.ctx, "failed to query logs", slog.Error(err))
|
||||
f.logger.Error(watchCtx, "failed to query logs", slog.Error(err))
|
||||
err = f.conn.Close(websocket.StatusInternalError, err.Error())
|
||||
if err != nil {
|
||||
f.logger.Warn(f.ctx, "failed to close websocket", slog.Error(err))
|
||||
f.logger.Warn(watchCtx, "failed to close websocket", slog.Error(err))
|
||||
}
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// Log the request immediately instead of after it completes.
|
||||
if rl := loggermw.RequestLoggerFromContext(f.ctx); rl != nil {
|
||||
rl.WriteLog(f.ctx, http.StatusAccepted)
|
||||
if rl := loggermw.RequestLoggerFromContext(watchCtx); rl != nil {
|
||||
rl.WriteLog(watchCtx, http.StatusAccepted)
|
||||
}
|
||||
|
||||
// no need to wait if the job is done
|
||||
@@ -614,14 +621,14 @@ func (f *logFollower) follow() {
|
||||
// We could soldier on and retry, but loss of database connectivity
|
||||
// is fairly serious, so instead just 500 and bail out. Client
|
||||
// can retry and hopefully find a healthier node.
|
||||
f.logger.Error(f.ctx, "dropped or corrupted notification", slog.Error(err))
|
||||
f.logger.Error(watchCtx, "dropped or corrupted notification", slog.Error(err))
|
||||
err = f.conn.Close(websocket.StatusInternalError, err.Error())
|
||||
if err != nil {
|
||||
f.logger.Warn(f.ctx, "failed to close websocket", slog.Error(err))
|
||||
f.logger.Warn(watchCtx, "failed to close websocket", slog.Error(err))
|
||||
}
|
||||
return
|
||||
case <-f.ctx.Done():
|
||||
// client disconnect
|
||||
case <-watchCtx.Done():
|
||||
// client disconnect or probe failure
|
||||
return
|
||||
case n := <-f.notifications:
|
||||
if n.EndOfLogs {
|
||||
@@ -630,14 +637,14 @@ func (f *logFollower) follow() {
|
||||
// gotten all logs prior to the start of our subscription.
|
||||
return
|
||||
}
|
||||
err = f.query()
|
||||
err = f.query(watchCtx)
|
||||
if err != nil {
|
||||
if f.ctx.Err() == nil && !xerrors.Is(err, io.EOF) {
|
||||
if watchCtx.Err() == nil && !xerrors.Is(err, io.EOF) {
|
||||
// neither context expiry, nor EOF, close and log
|
||||
f.logger.Error(f.ctx, "failed to query logs", slog.Error(err))
|
||||
f.logger.Error(watchCtx, "failed to query logs", slog.Error(err))
|
||||
err = f.conn.Close(websocket.StatusInternalError, httpapi.WebsocketCloseSprintf("%s", err.Error()))
|
||||
if err != nil {
|
||||
f.logger.Warn(f.ctx, "failed to close websocket", slog.Error(err))
|
||||
f.logger.Warn(watchCtx, "failed to close websocket", slog.Error(err))
|
||||
}
|
||||
}
|
||||
return
|
||||
@@ -673,9 +680,9 @@ func (f *logFollower) listener(_ context.Context, message []byte, err error) {
|
||||
|
||||
// query fetches the latest job logs from the database and writes them to the
|
||||
// connection.
|
||||
func (f *logFollower) query() error {
|
||||
f.logger.Debug(f.ctx, "querying logs", slog.F("after", f.after))
|
||||
logs, err := f.db.GetProvisionerLogsAfterID(f.ctx, database.GetProvisionerLogsAfterIDParams{
|
||||
func (f *logFollower) query(watchCtx context.Context) error {
|
||||
f.logger.Debug(watchCtx, "querying logs", slog.F("after", f.after))
|
||||
logs, err := f.db.GetProvisionerLogsAfterID(watchCtx, database.GetProvisionerLogsAfterIDParams{
|
||||
JobID: f.jobID,
|
||||
CreatedAfter: f.after,
|
||||
})
|
||||
@@ -688,7 +695,7 @@ func (f *logFollower) query() error {
|
||||
return xerrors.Errorf("error writing to websocket: %w", err)
|
||||
}
|
||||
f.after = log.ID
|
||||
f.logger.Debug(f.ctx, "wrote log to websocket", slog.F("id", log.ID))
|
||||
f.logger.Debug(watchCtx, "wrote log to websocket", slog.F("id", log.ID))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -19,11 +19,13 @@ import (
|
||||
"github.com/coder/coder/v2/coderd/database/dbmock"
|
||||
"github.com/coder/coder/v2/coderd/database/dbtime"
|
||||
"github.com/coder/coder/v2/coderd/database/pubsub"
|
||||
"github.com/coder/coder/v2/coderd/httpapi"
|
||||
"github.com/coder/coder/v2/coderd/httpmw/loggermw"
|
||||
"github.com/coder/coder/v2/coderd/httpmw/loggermw/loggermock"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/provisionersdk"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
"github.com/coder/quartz"
|
||||
"github.com/coder/websocket"
|
||||
)
|
||||
|
||||
@@ -150,6 +152,7 @@ func Test_logFollower_completeBeforeFollow(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
mDB := dbmock.NewMockStore(ctrl)
|
||||
ps := pubsub.NewInMemory()
|
||||
wsw := httpapi.NewWSWatcher(quartz.NewReal(), nil)
|
||||
now := dbtime.Now()
|
||||
job := database.ProvisionerJob{
|
||||
ID: uuid.New(),
|
||||
@@ -169,7 +172,7 @@ func Test_logFollower_completeBeforeFollow(t *testing.T) {
|
||||
|
||||
// we need an HTTP server to get a websocket
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(rw http.ResponseWriter, r *http.Request) {
|
||||
uut := newLogFollower(ctx, logger, mDB, ps, rw, r, job, 10)
|
||||
uut := newLogFollower(ctx, logger, mDB, ps, wsw, rw, r, job, 10)
|
||||
uut.follow()
|
||||
}))
|
||||
defer srv.Close()
|
||||
@@ -213,6 +216,7 @@ func Test_logFollower_completeBeforeSubscribe(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
mDB := dbmock.NewMockStore(ctrl)
|
||||
ps := pubsub.NewInMemory()
|
||||
wsw := httpapi.NewWSWatcher(quartz.NewReal(), nil)
|
||||
now := dbtime.Now()
|
||||
job := database.ProvisionerJob{
|
||||
ID: uuid.New(),
|
||||
@@ -230,7 +234,7 @@ func Test_logFollower_completeBeforeSubscribe(t *testing.T) {
|
||||
|
||||
// we need an HTTP server to get a websocket
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(rw http.ResponseWriter, r *http.Request) {
|
||||
uut := newLogFollower(ctx, logger, mDB, ps, rw, r, job, 0)
|
||||
uut := newLogFollower(ctx, logger, mDB, ps, wsw, rw, r, job, 0)
|
||||
uut.follow()
|
||||
}))
|
||||
defer srv.Close()
|
||||
@@ -291,6 +295,7 @@ func Test_logFollower_EndOfLogs(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
mDB := dbmock.NewMockStore(ctrl)
|
||||
ps := pubsub.NewInMemory()
|
||||
wsw := httpapi.NewWSWatcher(quartz.NewReal(), nil)
|
||||
now := dbtime.Now()
|
||||
job := database.ProvisionerJob{
|
||||
ID: uuid.New(),
|
||||
@@ -312,7 +317,7 @@ func Test_logFollower_EndOfLogs(t *testing.T) {
|
||||
|
||||
// we need an HTTP server to get a websocket
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(rw http.ResponseWriter, r *http.Request) {
|
||||
uut := newLogFollower(ctx, logger, mDB, ps, rw, r, job, 0)
|
||||
uut := newLogFollower(ctx, logger, mDB, ps, wsw, rw, r, job, 0)
|
||||
uut.follow()
|
||||
}))
|
||||
|
||||
|
||||
@@ -501,7 +501,7 @@ func (api *API) workspaceAgentLogs(rw http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
ctx, cancel := context.WithCancel(ctx)
|
||||
defer cancel()
|
||||
go httpapi.HeartbeatClose(ctx, api.Logger, cancel, conn)
|
||||
ctx = api.wsWatcher.Watch(ctx, api.Logger, conn)
|
||||
|
||||
encoder := wsjson.NewEncoder[[]codersdk.WorkspaceAgentLog](conn, websocket.MessageText)
|
||||
defer encoder.Close(websocket.StatusNormalClosure)
|
||||
@@ -861,7 +861,7 @@ func (api *API) watchWorkspaceAgentContainers(rw http.ResponseWriter, r *http.Re
|
||||
return
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithCancel(r.Context())
|
||||
ctx, cancel := context.WithCancel(ctx)
|
||||
defer cancel()
|
||||
|
||||
// Here we close the websocket for reading, so that the websocket library will handle pings and
|
||||
@@ -871,7 +871,7 @@ func (api *API) watchWorkspaceAgentContainers(rw http.ResponseWriter, r *http.Re
|
||||
ctx, wsNetConn := codersdk.WebsocketNetConn(ctx, conn, websocket.MessageText)
|
||||
defer wsNetConn.Close()
|
||||
|
||||
go httpapi.HeartbeatCloseWithClock(ctx, logger, cancel, conn, api.Clock)
|
||||
ctx = api.wsWatcher.Watch(ctx, logger, conn)
|
||||
|
||||
encoder := json.NewEncoder(wsNetConn)
|
||||
|
||||
@@ -1371,9 +1371,7 @@ func (api *API) workspaceAgentClientCoordinate(rw http.ResponseWriter, r *http.R
|
||||
ctx, wsNetConn := codersdk.WebsocketNetConn(ctx, conn, websocket.MessageBinary)
|
||||
defer wsNetConn.Close()
|
||||
|
||||
ctx, cancel := context.WithCancel(ctx)
|
||||
defer cancel()
|
||||
go httpapi.HeartbeatClose(ctx, api.Logger, cancel, conn)
|
||||
ctx = api.wsWatcher.Watch(ctx, api.Logger, conn)
|
||||
|
||||
defer conn.Close(websocket.StatusNormalClosure, "")
|
||||
err = api.TailnetClientService.ServeClient(ctx, version, wsNetConn, tailnet.StreamID{
|
||||
@@ -1670,7 +1668,7 @@ func (api *API) watchWorkspaceAgentMetadataSSE(rw http.ResponseWriter, r *http.R
|
||||
// @Router /api/v2/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.Logger))
|
||||
api.watchWorkspaceAgentMetadata(rw, r, httpapi.OneWayWebSocketEventSender(api.Logger, api.wsWatcher))
|
||||
}
|
||||
|
||||
func (api *API) watchWorkspaceAgentMetadata(
|
||||
@@ -2301,7 +2299,7 @@ func (api *API) tailnetRPCConn(rw http.ResponseWriter, r *http.Request) {
|
||||
|
||||
ctx, cancel := context.WithCancel(ctx)
|
||||
defer cancel()
|
||||
go httpapi.HeartbeatClose(ctx, api.Logger, cancel, conn)
|
||||
ctx = api.wsWatcher.Watch(ctx, api.Logger, conn)
|
||||
err = api.TailnetClientService.ServeClient(ctx, version, wsNetConn, tailnet.StreamID{
|
||||
Name: "client",
|
||||
ID: peerID,
|
||||
|
||||
@@ -134,6 +134,7 @@ func runWatchChatGitWorkspaceLookupTest(t *testing.T, workspaceErr error, wantSt
|
||||
Authorizer: &mockAuthorizer{},
|
||||
Logger: logger,
|
||||
},
|
||||
wsWatcher: httpapi.NewWSWatcher(quartz.NewReal(), nil),
|
||||
}
|
||||
)
|
||||
|
||||
@@ -190,6 +191,7 @@ func TestWatchChatGit(t *testing.T) {
|
||||
Logger: logger,
|
||||
DeploymentValues: &codersdk.DeploymentValues{},
|
||||
},
|
||||
wsWatcher: httpapi.NewWSWatcher(quartz.NewReal(), nil),
|
||||
}
|
||||
)
|
||||
|
||||
@@ -264,6 +266,7 @@ func TestWatchChatGit(t *testing.T) {
|
||||
Logger: logger,
|
||||
DeploymentValues: &codersdk.DeploymentValues{},
|
||||
},
|
||||
wsWatcher: httpapi.NewWSWatcher(quartz.NewReal(), nil),
|
||||
}
|
||||
)
|
||||
|
||||
@@ -424,6 +427,7 @@ func TestWatchChatGit(t *testing.T) {
|
||||
Authorizer: &mockAuthorizer{},
|
||||
Logger: logger,
|
||||
},
|
||||
wsWatcher: httpapi.NewWSWatcher(quartz.NewReal(), nil),
|
||||
}
|
||||
)
|
||||
|
||||
@@ -602,6 +606,7 @@ func TestWatchChatGit(t *testing.T) {
|
||||
Authorizer: &mockAuthorizer{},
|
||||
Logger: logger,
|
||||
},
|
||||
wsWatcher: httpapi.NewWSWatcher(quartz.NewReal(), nil),
|
||||
}
|
||||
)
|
||||
|
||||
@@ -773,10 +778,12 @@ func TestWatchAgentContainers(t *testing.T) {
|
||||
DeploymentValues: &codersdk.DeploymentValues{},
|
||||
TailnetCoordinator: tailnettest.NewFakeCoordinator(),
|
||||
},
|
||||
wsWatcher: httpapi.NewWSWatcher(mClock, nil),
|
||||
}
|
||||
)
|
||||
|
||||
trap := mClock.Trap().NewTicker("HeartbeatClose")
|
||||
trap := mClock.Trap().NewTicker("WSWatcher")
|
||||
|
||||
defer trap.Close()
|
||||
|
||||
var tailnetCoordinator tailnet.Coordinator = mCoordinator
|
||||
@@ -897,6 +904,7 @@ func TestWatchAgentContainers(t *testing.T) {
|
||||
DeploymentValues: &codersdk.DeploymentValues{},
|
||||
TailnetCoordinator: tailnettest.NewFakeCoordinator(),
|
||||
},
|
||||
wsWatcher: httpapi.NewWSWatcher(quartz.NewReal(), nil),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@@ -112,6 +112,7 @@ type ServerOptions struct {
|
||||
|
||||
AgentProvider AgentProvider
|
||||
StatsCollector *StatsCollector
|
||||
WSWatcher *httpapi.WSWatcher
|
||||
}
|
||||
|
||||
// Server serves workspace apps endpoints, including:
|
||||
@@ -765,11 +766,12 @@ func (s *Server) workspaceAgentPTY(rw http.ResponseWriter, r *http.Request) {
|
||||
})
|
||||
return
|
||||
}
|
||||
go httpapi.HeartbeatClose(ctx, s.Logger, cancel, conn)
|
||||
|
||||
ctx, wsNetConn := WebsocketNetConn(ctx, conn, websocket.MessageBinary)
|
||||
defer wsNetConn.Close() // Also closes conn.
|
||||
|
||||
ctx = s.WSWatcher.Watch(ctx, s.Logger, conn)
|
||||
|
||||
agentConn, release, err := s.AgentProvider.AgentConn(ctx, appToken.AgentID)
|
||||
if err != nil {
|
||||
log.Debug(ctx, "dial workspace agent", slog.Error(err))
|
||||
|
||||
@@ -2033,7 +2033,7 @@ func (api *API) watchWorkspaceSSE(rw http.ResponseWriter, r *http.Request) {
|
||||
// @Success 200 {object} codersdk.ServerSentEvent
|
||||
// @Router /api/v2/workspaces/{workspace}/watch-ws [get]
|
||||
func (api *API) watchWorkspaceWS(rw http.ResponseWriter, r *http.Request) {
|
||||
api.watchWorkspace(rw, r, httpapi.OneWayWebSocketEventSender(api.Logger))
|
||||
api.watchWorkspace(rw, r, httpapi.OneWayWebSocketEventSender(api.Logger, api.wsWatcher))
|
||||
}
|
||||
|
||||
func (api *API) watchWorkspace(
|
||||
@@ -2230,7 +2230,7 @@ func (api *API) watchAllWorkspaceBuilds(rw http.ResponseWriter, r *http.Request)
|
||||
_ = conn.CloseRead(context.Background())
|
||||
|
||||
ctx, cancel := context.WithCancel(ctx)
|
||||
go httpapi.HeartbeatClose(ctx, api.Logger, cancel, conn)
|
||||
ctx = api.wsWatcher.Watch(ctx, api.Logger, conn)
|
||||
defer cancel()
|
||||
|
||||
enc := wsjson.NewEncoder[codersdk.WorkspaceBuildUpdate](conn, websocket.MessageText)
|
||||
|
||||
Reference in New Issue
Block a user