Hugo Dutka
2026-06-12 13:33:12 +02:00
committed by GitHub
parent 4a07f61c50
commit 4debd23cbb
155 changed files with 37612 additions and 24513 deletions
+9 -9
View File
@@ -159,10 +159,14 @@ func New(ctx context.Context, options *Options) (_ *API, err error) {
}
var replicaManagerPtr atomic.Pointer[replicasync.Manager]
var api *API
resolveReplicaAddress := func(
_ context.Context,
replicaID uuid.UUID,
) (string, bool) {
if api != nil && api.AGPL != nil && replicaID == api.AGPL.ID && api.AGPL.AccessURL != nil {
return api.AGPL.AccessURL.String(), true
}
manager := replicaManagerPtr.Load()
if manager == nil {
return "", false
@@ -180,7 +184,7 @@ func New(ctx context.Context, options *Options) (_ *API, err error) {
return "", false
}
api := &API{
api = &API{
ctx: ctx,
cancel: cancelFunc,
Options: options,
@@ -207,17 +211,13 @@ func New(ctx context.Context, options *Options) (_ *API, err error) {
replicaHTTPClient = http.DefaultClient
}
// Use a closure that captures api by reference so it can access
// api.AGPL.ID after coderd.New is called. The SubscribeFn is
// only invoked from Subscribe, which happens after init.
options.Options.ChatSubscribeFn = entchatd.NewMultiReplicaSubscribeFn(entchatd.MultiReplicaSubscribeConfig{
// api.AGPL.ID after coderd.New is called. The parts dialer is
// only invoked from stream subscriptions, which happen after init.
options.Options.ChatStreamPartsDialer = entchatd.NewStreamPartsDialer(entchatd.StreamPartsDialerConfig{
ResolveReplicaAddress: resolveReplicaAddress,
ReplicaHTTPClient: replicaHTTPClient,
ReplicaIDFn: func() uuid.UUID {
id := api.AGPL.ID
if id == uuid.Nil {
return uuid.New()
}
return id
return api.AGPL.ID
},
})
+43 -778
View File
@@ -2,25 +2,16 @@ package chatd
import (
"context"
"errors"
"fmt"
"net/http"
"net/url"
"strconv"
"strings"
"time"
"github.com/google/uuid"
"golang.org/x/xerrors"
"cdr.dev/slog/v3"
"github.com/coder/coder/v2/coderd/database"
osschatd "github.com/coder/coder/v2/coderd/x/chatd"
"github.com/coder/coder/v2/codersdk"
"github.com/coder/quartz"
"github.com/coder/retry"
"github.com/coder/websocket"
"github.com/coder/websocket/wsjson"
)
// RelaySourceHeader marks replica-relayed stream requests.
@@ -29,28 +20,10 @@ const RelaySourceHeader = "X-Coder-Relay-Source-Replica"
const (
authorizationHeader = "Authorization"
cookieHeader = "Cookie"
// relayDrainTimeout is how long an established relay is
// kept open after the chat leaves running state, giving
// buffered snapshot events time to be forwarded before
// the relay is torn down.
relayDrainTimeout = 200 * time.Millisecond
// Retry knobs for the cross-replica relay handshake. Uses the
// github.com/coder/retry defaults (φ-growth, no jitter) but drives
// the delay manually because retry.Retrier.Wait uses time.After,
// which isn't compatible with quartz.Clock determinism in tests.
relayRetryFloor = 500 * time.Millisecond // first retry matches old fixed delay
relayRetryCeil = 15 * time.Second // cap stall before tear-down
// After this many reconnect retries the relay leg is torn down.
// Total dial attempts = 1 initial dial + relayMaxRetries.
relayMaxRetries = 6
)
// RelayDialError wraps a failed relay handshake. HTTPStatus is 0
// when the failure happened before a response (DNS, TCP, TLS,
// timeout, context cancel); otherwise it carries the peer's status
// code for the reconnect loop to classify.
// when the failure happened before a response.
type RelayDialError struct {
HTTPStatus int
Err error
@@ -60,661 +33,59 @@ func (e *RelayDialError) Error() string { return e.Err.Error() }
func (e *RelayDialError) Unwrap() error { return e.Err }
// IsUnrecoverable reports whether retrying with the same captured
// session token is futile. Only 401/403 qualify - the token is dead
// or the peer won't authorize it. 5xx, 429, network, and context
// errors fall through to backoff.
// session token is futile.
func (e *RelayDialError) IsUnrecoverable() bool {
return e.HTTPStatus == http.StatusUnauthorized ||
e.HTTPStatus == http.StatusForbidden
}
// MultiReplicaSubscribeConfig holds the dependencies for multi-replica chat
// subscription. ReplicaIDFn is called lazily because the
// replica ID may not be known at construction time.
//
// DialerFn, when set, overrides the default WebSocket relay
// dialer. This is used in tests to inject mock relay behavior
// without requiring real HTTP servers.
type MultiReplicaSubscribeConfig struct {
// StreamPartsDialerConfig holds dependencies for multi-replica stream parts.
type StreamPartsDialerConfig struct {
ResolveReplicaAddress func(context.Context, uuid.UUID) (string, bool)
ReplicaHTTPClient *http.Client
ReplicaIDFn func() uuid.UUID
DialerFn func(
ctx context.Context,
chatID uuid.UUID,
workerID uuid.UUID,
requestHeader http.Header,
) (
snapshot []codersdk.ChatStreamEvent,
parts <-chan codersdk.ChatStreamEvent,
cancel func(),
err error,
)
// Clock is used for creating timers. In production use
// quartz.NewReal(); in tests use quartz.NewMock(t) to
// control reconnect timing deterministically.
Clock quartz.Clock
DialerFn func(context.Context, osschatd.StreamPartsDialInput) (osschatd.StreamPartsSession, error)
}
// dial returns the configured dialer, preferring DialerFn (tests)
// over the real dialRelay. Returns nil when relay is not configured.
func (c MultiReplicaSubscribeConfig) dial() func(
// NewStreamPartsDialer returns a dialer for the owning replica's parts endpoint.
func NewStreamPartsDialer(cfg StreamPartsDialerConfig) osschatd.StreamPartsDialer {
return func(ctx context.Context, input osschatd.StreamPartsDialInput) (osschatd.StreamPartsSession, error) {
if cfg.DialerFn != nil {
return cfg.DialerFn(ctx, input)
}
return dialRelayParts(ctx, input, cfg)
}
}
func dialRelayParts(
ctx context.Context,
chatID uuid.UUID,
workerID uuid.UUID,
requestHeader http.Header,
) (
[]codersdk.ChatStreamEvent,
<-chan codersdk.ChatStreamEvent,
func(),
error,
) {
if c.DialerFn != nil {
return c.DialerFn
input osschatd.StreamPartsDialInput,
cfg StreamPartsDialerConfig,
) (osschatd.StreamPartsSession, error) {
if cfg.ResolveReplicaAddress == nil {
return nil, &RelayDialError{Err: xerrors.New("dial relay stream parts: resolver not configured")}
}
if c.ResolveReplicaAddress == nil {
return nil
}
return func(
ctx context.Context,
chatID uuid.UUID,
workerID uuid.UUID,
requestHeader http.Header,
) (
[]codersdk.ChatStreamEvent,
<-chan codersdk.ChatStreamEvent,
func(),
error,
) {
return dialRelay(ctx, chatID, workerID, requestHeader, c, c.clock())
}
}
// clock returns the quartz.Clock to use. Defaults to a real clock
// when not set.
func (c MultiReplicaSubscribeConfig) clock() quartz.Clock {
if c.Clock != nil {
return c.Clock
}
return quartz.NewReal()
}
// NewMultiReplicaSubscribeFn returns a SubscribeFn that manages
// relay connections to remote replicas and returns relay
// message_part events only. OSS handles pubsub subscription,
// message catch-up, queue updates, status forwarding, and local
// parts merging.
//
//nolint:gocognit // Complexity is inherent to the multi-source merge loop.
func NewMultiReplicaSubscribeFn(
cfg MultiReplicaSubscribeConfig,
) osschatd.SubscribeFn {
return func(ctx context.Context, params osschatd.SubscribeFnParams) <-chan codersdk.ChatStreamEvent {
chatID := params.ChatID
requestHeader := params.RequestHeader
logger := params.Logger
var relayCancel func()
var relayParts <-chan codersdk.ChatStreamEvent
// If the chat is currently running on a different worker
// and we have a remote parts provider, open an initial
// relay synchronously so the caller gets in-flight
// message_part events right away.
var initialRelaySnapshot []codersdk.ChatStreamEvent
if params.Chat.Status == database.ChatStatusRunning &&
params.Chat.WorkerID.Valid &&
params.Chat.WorkerID.UUID != params.WorkerID &&
cfg.dial() != nil {
snapshot, parts, cancel, err := cfg.dial()(ctx, chatID, params.Chat.WorkerID.UUID, requestHeader)
if err == nil {
relayCancel = cancel
relayParts = parts
// Collect relay message_parts to forward at the
// start of the merge goroutine.
for _, event := range snapshot {
if event.Type == codersdk.ChatStreamEventTypeMessagePart {
initialRelaySnapshot = append(initialRelaySnapshot, event)
}
}
} else {
logger.Warn(ctx, "failed to open initial relay for chat stream",
slog.F("chat_id", chatID),
slog.Error(err),
)
}
}
// Merge all event sources.
mergedEvents := make(chan codersdk.ChatStreamEvent, 128)
// Channel for async relay establishment.
type relayResult struct {
parts <-chan codersdk.ChatStreamEvent
cancel func()
workerID uuid.UUID // the worker this dial targeted
// err and parts are mutually exclusive: success sets
// parts; failure sets err (unwrap to *RelayDialError
// for classification).
err error
}
relayReadyCh := make(chan relayResult, 4)
// Reset on successful dial or when the relay target
// changes, so a fresh target starts at the floor delay.
retryState := newRelayRetryState()
// Per-dial context so in-flight dials can be canceled when
// a new dial is initiated or the relay is closed.
var dialCancel context.CancelFunc
// expectedWorkerID tracks which replica we expect the next
// relay result to target. Stale results are discarded.
var expectedWorkerID uuid.UUID
// Reconnect timer state.
var reconnectTimer *quartz.Timer
var reconnectCh <-chan time.Time
// drainAndClose is set when the chat transitions away
// from running while a relay dial is still in progress.
// Instead of canceling the dial immediately, we let it
// complete so the snapshot of buffered message_parts
// can be forwarded to the subscriber.
var drainAndClose bool
// Drain timer state. When the relay connects in
// drain-and-close mode, a short timer is started.
// During this window the normal relayPartsCh case
// forwards buffered snapshot events. When the timer
// fires the relay is torn down.
var drainTimer *quartz.Timer
var drainTimerCh <-chan time.Time
// Helper to close relay and stop any pending reconnect
// timer.
closeRelay := func() {
// Cancel any in-flight dial goroutine first.
if dialCancel != nil {
dialCancel()
dialCancel = nil
}
// Drain all buffered relay results from canceled dials.
for {
select {
case result := <-relayReadyCh:
if result.cancel != nil {
result.cancel()
}
default:
goto drained
}
}
drained:
expectedWorkerID = uuid.Nil
if relayCancel != nil {
relayCancel()
relayCancel = nil
}
relayParts = nil
if reconnectTimer != nil {
reconnectTimer.Stop()
reconnectTimer = nil
reconnectCh = nil
}
if drainTimer != nil {
drainTimer.Stop()
drainTimer = nil
drainTimerCh = nil
}
drainAndClose = false
}
// openRelayAsync dials the remote replica in a background
// goroutine and delivers the result on relayReadyCh so the
// main select loop is never blocked by network I/O.
openRelayAsync := func(workerID uuid.UUID) {
if cfg.dial() == nil {
return
}
// Scoped here (not in closeRelay) so repeated dials
// against the same worker keep the attempt counter and
// correctly trip the cap.
if workerID != expectedWorkerID {
retryState.reset()
}
closeRelay()
// Create a per-dial context so this goroutine is
// canceled if closeRelay() or openRelayAsync() is
// called again before the dial completes.
var dialCtx context.Context
dialCtx, dialCancel = context.WithCancel(ctx)
expectedWorkerID = workerID
go func() {
snapshot, parts, cancel, err := cfg.dial()(dialCtx, chatID, workerID, requestHeader)
if err != nil {
// Don't log context-canceled errors
// since they are expected when a dial is
// superseded by a newer one.
if dialCtx.Err() == nil {
fields := []slog.Field{
slog.F("chat_id", chatID),
slog.F("worker_id", workerID),
slog.Error(err),
}
// Surface the peer's HTTP status (when we
// got one) as a structured field so
// operators can filter 401/403 spam
// separately from 5xx/network warnings.
var dialErr *RelayDialError
if errors.As(err, &dialErr) && dialErr.HTTPStatus != 0 {
fields = append(fields, slog.F("http_status", dialErr.HTTPStatus))
}
logger.Warn(ctx, "failed to open relay for message parts", fields...)
}
// Hand the error to the merge loop, which will
// classify it and either back off or tear down.
select {
case relayReadyCh <- relayResult{workerID: workerID, err: err}:
case <-dialCtx.Done():
}
return
}
// Discard stale dials so we don't start a
// wrappedParts goroutine on a canceled connection.
if dialCtx.Err() != nil {
cancel()
return
}
// Wrap the relay channel so snapshot parts
// are delivered through the same channel as
// live parts. This goroutine only forwards
// events - it does not own the relay
// lifecycle. When dialCtx is canceled it
// simply returns, closing wrappedParts via
// its defer. The cancel() is called by
// whoever canceled dialCtx (closeRelay or
// the send-fallback select below).
wrappedParts := make(chan codersdk.ChatStreamEvent, 128)
go func() {
defer close(wrappedParts)
for _, event := range snapshot {
if event.Type == codersdk.ChatStreamEventTypeMessagePart {
select {
case wrappedParts <- event:
case <-dialCtx.Done():
return
}
}
}
for {
select {
case event, ok := <-parts:
if !ok {
return
}
select {
case wrappedParts <- event:
case <-dialCtx.Done():
return
}
case <-dialCtx.Done():
return
}
}
}()
select {
case relayReadyCh <- relayResult{parts: wrappedParts, cancel: cancel, workerID: workerID}:
case <-dialCtx.Done():
cancel()
}
}()
}
// scheduleRelayReconnect arms a timer so the select loop
// can re-check chat status and reopen the relay. Callers
// pass the delay from retryState so the failed-dial branch
// gets backoff while transient branches stay at the floor.
scheduleRelayReconnect := func(delay time.Duration) {
if cfg.dial() == nil {
return
}
if reconnectTimer != nil {
reconnectTimer.Stop()
}
reconnectTimer = cfg.clock().NewTimer(delay, "reconnect")
reconnectCh = reconnectTimer.C
}
// sendRelayTerminalError enqueues one error event for the
// subscriber; callers return afterwards so the deferred
// close(mergedEvents) fires and the OSS merge loop tears
// the relay leg down while pubsub/local sources keep going.
sendRelayTerminalError := func(msg string) {
select {
case mergedEvents <- codersdk.ChatStreamEvent{
Type: codersdk.ChatStreamEventTypeError,
ChatID: chatID,
Error: &codersdk.ChatError{Message: msg},
}:
case <-ctx.Done():
}
}
statusNotifications := params.StatusNotifications
go func() {
defer close(mergedEvents)
defer closeRelay()
// Forward any initial relay snapshot parts
// collected synchronously above.
for _, event := range initialRelaySnapshot {
select {
case <-ctx.Done():
return
case mergedEvents <- event:
}
}
for {
relayPartsCh := relayParts
select {
case <-ctx.Done():
return
case result := <-relayReadyCh:
// Discard stale relay results from a
// previous dial that was superseded.
if result.workerID != expectedWorkerID {
if result.cancel != nil {
result.cancel()
}
continue
}
// A nil parts channel signals the dial
// failed - classify the error to decide
// whether to schedule a backoff retry, emit a
// terminal error and tear the relay leg down
// (unrecoverable / cap reached), or simply
// drop the stale drain.
if result.parts == nil {
if drainAndClose {
// Dial failed and we were only
// waiting to drain - nothing to do.
drainAndClose = false
continue
}
var dialErr *RelayDialError
if errors.As(result.err, &dialErr) && dialErr.IsUnrecoverable() {
logger.Warn(ctx, "relay dial unrecoverable; tearing down relay leg",
slog.F("chat_id", chatID),
slog.F("worker_id", result.workerID),
slog.F("http_status", dialErr.HTTPStatus),
)
sendRelayTerminalError(fmt.Sprintf(
"relay authentication failed (status %d)",
dialErr.HTTPStatus,
))
return
}
delay, giveUp := retryState.next()
if giveUp {
logger.Warn(ctx, "relay dial retry cap reached; tearing down relay leg",
slog.F("chat_id", chatID),
slog.F("worker_id", result.workerID),
slog.F("max_retries", relayMaxRetries),
)
sendRelayTerminalError(fmt.Sprintf(
"relay connection failed after %d retries",
relayMaxRetries,
))
return
}
scheduleRelayReconnect(delay)
continue
}
// An async relay dial completed. Swap in the
// new relay channel. We deliberately do NOT
// reset the retry counter here: a peer that
// accepts the handshake and immediately drops
// the stream would otherwise keep reconnecting
// forever, since each success would zero the
// counter before the next drop re-incremented
// it. The counter only resets when the target
// worker changes (see openRelayAsync).
if relayCancel != nil {
relayCancel()
relayCancel = nil
}
relayParts = result.parts
relayCancel = result.cancel
if drainAndClose {
// The chat is no longer running on
// the remote worker, but the dial
// completed. Verify no new worker
// has claimed the chat before we
// drain stale parts.
currentChat, dbErr := params.DB.GetChatByID(ctx, chatID)
if dbErr != nil {
logger.Warn(ctx, "failed to check chat status for relay drain",
slog.F("chat_id", chatID),
slog.Error(dbErr),
)
}
if dbErr == nil && currentChat.Status == database.ChatStatusRunning &&
currentChat.WorkerID.Valid &&
currentChat.WorkerID.UUID != params.WorkerID {
// A new worker picked up the chat;
// discard the stale relay and let
// openRelayAsync handle the new one.
closeRelay()
} else {
// Chat is still idle - drain the
// buffered snapshot before closing.
if drainTimer != nil {
drainTimer.Stop()
}
drainTimer = cfg.clock().NewTimer(relayDrainTimeout, "drain")
drainTimerCh = drainTimer.C
drainAndClose = false
}
}
case <-reconnectCh:
reconnectCh = nil
// Re-check whether the chat is still
// running on a remote worker before
// reconnecting.
currentChat, chatErr := params.DB.GetChatByID(ctx, chatID)
if chatErr != nil {
logger.Warn(ctx, "failed to get chat for relay reconnect",
slog.F("chat_id", chatID),
slog.Error(chatErr),
)
// Retry on transient DB errors to
// avoid permanently stalling the
// stream. The same retry state
// bounds the DB-error loop too so a
// persistently broken DB eventually
// tears the relay down instead of
// spinning forever.
delay, giveUp := retryState.next()
if giveUp {
logger.Warn(ctx, "relay reconnect retry cap reached; tearing down relay leg",
slog.F("chat_id", chatID),
slog.F("max_retries", relayMaxRetries),
)
sendRelayTerminalError(fmt.Sprintf(
"relay connection failed after %d retries",
relayMaxRetries,
))
return
}
scheduleRelayReconnect(delay)
continue
}
if currentChat.Status == database.ChatStatusRunning &&
currentChat.WorkerID.Valid && currentChat.WorkerID.UUID != params.WorkerID {
openRelayAsync(currentChat.WorkerID.UUID)
}
case sn, ok := <-statusNotifications:
if !ok {
statusNotifications = nil
continue
}
if sn.Status == database.ChatStatusRunning && sn.WorkerID != uuid.Nil && sn.WorkerID != params.WorkerID {
openRelayAsync(sn.WorkerID)
} else {
switch {
case dialCancel != nil && relayParts == nil:
// In-progress dial: let it complete
// so its snapshot can be forwarded.
drainAndClose = true
case relayParts != nil:
// Active relay: give it a short
// window to deliver any remaining
// buffered parts before closing.
if drainTimer != nil {
drainTimer.Stop()
}
drainTimer = cfg.clock().NewTimer(relayDrainTimeout, "drain")
drainTimerCh = drainTimer.C
default:
closeRelay()
}
}
case <-drainTimerCh:
drainTimerCh = nil
drainTimer = nil
closeRelay()
case event, ok := <-relayPartsCh:
if !ok {
if relayCancel != nil {
relayCancel()
relayCancel = nil
}
relayParts = nil
// Reuse the retry state so a relay that
// repeatedly drops eventually tears down.
delay, giveUp := retryState.next()
if giveUp {
logger.Warn(ctx, "relay drop retry cap reached; tearing down relay leg",
slog.F("chat_id", chatID),
slog.F("max_retries", relayMaxRetries),
)
sendRelayTerminalError(fmt.Sprintf(
"relay connection failed after %d retries",
relayMaxRetries,
))
return
}
scheduleRelayReconnect(delay)
continue
}
// Only forward message_part events from
// relay.
if event.Type == codersdk.ChatStreamEventTypeMessagePart {
select {
case <-ctx.Done():
return
case mergedEvents <- event:
}
}
}
}
}()
// Cleanup is driven by ctx cancellation: the merge
// goroutine owns all relay state (reconnectTimer,
// relayCancel, dialCancel, etc.) and tears it down
// via defer closeRelay() when ctx is done.
return mergedEvents
}
}
// relayRetryState drives the retry policy for the relay reconnect
// loop. Wraps github.com/coder/retry to reuse its φ-growth defaults
// but computes the delay without blocking so the merge loop can
// schedule its own quartz.Clock timer.
//
// Not safe for concurrent use.
type relayRetryState struct {
retrier *retry.Retrier
attempts int
}
func newRelayRetryState() *relayRetryState {
return &relayRetryState{
retrier: retry.New(relayRetryFloor, relayRetryCeil),
}
}
// next returns the delay before the next dial and sets giveUp once
// attempts exceed relayMaxRetries. Adapts the math from
// retry.Retrier.Wait (github.com/coder/retry/retrier.go) without
// blocking: the library's Wait returns 0 on the first call and sets
// Delay to Floor only after the sleep, so we clamp to Floor up
// front.
func (s *relayRetryState) next() (delay time.Duration, giveUp bool) {
s.attempts++
if s.attempts > relayMaxRetries {
return 0, true
}
r := s.retrier
d := time.Duration(float64(r.Delay) * r.Rate)
if d > r.Ceil {
d = r.Ceil
}
if d < r.Floor {
d = r.Floor
}
r.Delay = d
return d, false
}
// reset returns the state to the floor delay and zero attempts.
// Called after a successful dial or a relay target change.
func (s *relayRetryState) reset() {
s.retrier.Reset()
s.attempts = 0
}
// dialRelay opens a WebSocket to the replica owning chatID and
// returns any buffered message_part snapshot plus a live channel of
// subsequent events. Handshake failures return an error unwrapping
// to *RelayDialError so callers can classify via IsUnrecoverable.
//
// websocket.Dial is called directly (not via the SDK wrapper) so we
// can read *http.Response.StatusCode for classification.
func dialRelay(
ctx context.Context,
chatID uuid.UUID,
workerID uuid.UUID,
requestHeader http.Header,
cfg MultiReplicaSubscribeConfig,
clk quartz.Clock,
) (
snapshot []codersdk.ChatStreamEvent,
parts <-chan codersdk.ChatStreamEvent,
cancel func(),
err error,
) {
address, ok := cfg.ResolveReplicaAddress(ctx, workerID)
address, ok := cfg.ResolveReplicaAddress(ctx, input.WorkerID)
if !ok {
return nil, nil, nil, &RelayDialError{
Err: xerrors.New("dial relay stream: worker replica not found"),
}
return nil, &RelayDialError{Err: xerrors.New("dial relay stream parts: worker replica not found")}
}
wsURL, err := buildRelayURL(address, chatID)
wsURL, err := buildRelayURL(address, input.ChatID)
if err != nil {
return nil, nil, nil, &RelayDialError{
Err: xerrors.Errorf("dial relay stream: %w", err),
}
return nil, &RelayDialError{Err: xerrors.Errorf("dial relay stream parts: %w", err)}
}
if cfg.ReplicaIDFn == nil {
return nil, &RelayDialError{Err: xerrors.New("dial relay stream parts: replica ID function not configured")}
}
replicaID := cfg.ReplicaIDFn()
if replicaID == uuid.Nil {
return nil, &RelayDialError{Err: xerrors.New("dial relay stream parts: replica ID is nil")}
}
headers := make(http.Header, 2)
headers.Set(codersdk.SessionTokenHeader, extractSessionToken(requestHeader))
headers.Set(codersdk.SessionTokenHeader, extractSessionToken(input.RequestHeader))
headers.Set(RelaySourceHeader, replicaID.String())
relayCtx, relayCancel := context.WithCancel(ctx)
conn, resp, dialErr := websocket.Dial(relayCtx, wsURL, &websocket.DialOptions{
conn, resp, dialErr := websocket.Dial(ctx, wsURL, &websocket.DialOptions{
HTTPClient: cfg.ReplicaHTTPClient,
HTTPHeader: headers,
CompressionMode: websocket.CompressionDisabled,
@@ -722,118 +93,22 @@ func dialRelay(
status := 0
if resp != nil {
status = resp.StatusCode
// The websocket library closes resp.Body on success; on
// failure we close it ourselves so we don't leak the TCP
// connection.
if dialErr != nil && resp.Body != nil {
_ = resp.Body.Close()
}
}
if dialErr != nil {
relayCancel()
return nil, nil, nil, &RelayDialError{
return nil, &RelayDialError{
HTTPStatus: status,
Err: xerrors.Errorf("dial relay stream: %w", dialErr),
Err: xerrors.Errorf("dial relay stream parts: %w", dialErr),
}
}
// Match the server's 4 MiB read limit in codersdk.StreamChat so
// large message_part batches don't trip the default 32 KiB cap.
conn.SetReadLimit(1 << 22)
snapshot = make([]codersdk.ChatStreamEvent, 0, 100)
// sourceEvents is the flattened batch→event channel. A small
// goroutine reads batches off the websocket and fans them out;
// callers see a single event stream identical to the shape the
// old SDK call produced.
sourceEvents := make(chan codersdk.ChatStreamEvent, 128)
go func() {
defer close(sourceEvents)
for {
var batch []codersdk.ChatStreamEvent
if readErr := wsjson.Read(relayCtx, conn, &batch); readErr != nil {
return
}
for _, event := range batch {
select {
case sourceEvents <- event:
case <-relayCtx.Done():
return
}
}
}
}()
closeSource := func() {
relayCancel()
_ = conn.Close(websocket.StatusNormalClosure, "")
}
// Wait briefly for the first event to handle the common
// case where the remote side has buffered parts but hasn't
// flushed them to the WebSocket yet.
const drainTimeout = time.Second
drainTimer := clk.NewTimer(drainTimeout, "drain")
defer drainTimer.Stop()
drainInitial:
for len(snapshot) < cap(snapshot) {
select {
case <-relayCtx.Done():
closeSource()
return nil, nil, nil, &RelayDialError{
Err: xerrors.Errorf("dial relay stream: %w", relayCtx.Err()),
}
case event, ok := <-sourceEvents:
if !ok {
break drainInitial
}
if event.Type != codersdk.ChatStreamEventTypeMessagePart {
continue
}
snapshot = append(snapshot, event)
// After getting the first event, switch to
// non-blocking drain for remaining buffered events.
drainTimer.Stop()
drainTimer.Reset(0)
case <-drainTimer.C:
break drainInitial
}
}
events := make(chan codersdk.ChatStreamEvent, 128)
go func() {
defer close(events)
defer closeSource()
// No need to re-send snapshot events - they're
// returned to the caller directly.
for {
select {
case <-relayCtx.Done():
return
case event, ok := <-sourceEvents:
if !ok {
return
}
if event.Type != codersdk.ChatStreamEventTypeMessagePart {
continue
}
select {
case events <- event:
case <-relayCtx.Done():
return
}
}
}
}()
return snapshot, events, closeSource, nil
return osschatd.NewStreamPartsJSONSession(ctx, conn), nil
}
// buildRelayURL builds the websocket URL for the chat stream
// endpoint on a peer replica. It maps http(s) schemes to ws(s).
// buildRelayURL builds the websocket URL for the chat stream parts endpoint on
// a peer replica. It maps http(s) schemes to ws(s).
func buildRelayURL(address string, chatID uuid.UUID) (string, error) {
u, err := url.Parse(address)
if err != nil {
@@ -845,40 +120,30 @@ func buildRelayURL(address string, chatID uuid.UUID) (string, error) {
case "https":
u.Scheme = "wss"
case "ws", "wss":
// already a websocket URL, leave as-is.
default:
return "", xerrors.Errorf("unsupported relay address scheme %q", u.Scheme)
}
u.Path = fmt.Sprintf("/api/experimental/chats/%s/stream", chatID)
q := u.Query()
// Relays only need live message_part events, not the full
// history; pass the relay sentinel so the peer skips its
// durable DB snapshot and delivers in-flight parts only.
q.Set("after_id", strconv.FormatInt(osschatd.RelaySentinelAfterID, 10))
u.RawQuery = q.Encode()
u.Path = "/api/experimental/chats/" + chatID.String() + "/stream/parts"
u.RawQuery = ""
return u.String(), nil
}
// extractSessionToken returns the session token carried by the
// given request headers. It mirrors the priority order used by
// apiKeyMiddleware: cookie, then Coder-Session-Token header, then
// Authorization: Bearer header.
// extractSessionToken returns the session token carried by the given request
// headers. It mirrors the priority order used by apiKeyMiddleware: cookie,
// then Coder-Session-Token header, then Authorization: Bearer header.
func extractSessionToken(header http.Header) string {
if header == nil {
return ""
}
// Cookie (browser WebSocket upgrade - most common relay case).
if raw := header.Get(cookieHeader); raw != "" {
r := &http.Request{Header: http.Header{cookieHeader: {raw}}}
if c, err := r.Cookie(codersdk.SessionTokenCookie); err == nil && c.Value != "" {
return c.Value
}
}
// Coder-Session-Token header (SDK / CLI callers).
if v := header.Get(codersdk.SessionTokenHeader); v != "" {
return v
}
// Authorization: Bearer <token>.
if v := header.Get(authorizationHeader); len(v) > 7 && strings.EqualFold(v[:7], "bearer ") {
return strings.TrimSpace(v[7:])
}
@@ -1,796 +0,0 @@
package chatd_test
import (
"context"
"database/sql"
"encoding/json"
"io"
"math"
"net/http"
"net/http/httptest"
"regexp"
"sync/atomic"
"testing"
"time"
"github.com/google/uuid"
"github.com/stretchr/testify/require"
"golang.org/x/xerrors"
"cdr.dev/slog/v3/sloggers/slogtest"
"github.com/coder/coder/v2/coderd/database"
"github.com/coder/coder/v2/coderd/database/dbtestutil"
dbpubsub "github.com/coder/coder/v2/coderd/database/pubsub"
coderdpubsub "github.com/coder/coder/v2/coderd/pubsub"
osschatd "github.com/coder/coder/v2/coderd/x/chatd"
"github.com/coder/coder/v2/codersdk"
entchatd "github.com/coder/coder/v2/enterprise/coderd/x/chatd"
"github.com/coder/coder/v2/testutil"
"github.com/coder/quartz"
)
// mulPhi multiplies a duration by math.Phi to compute the next
// step in retry.Retrier's φ-growth backoff sequence. If
// TestRelayReconnectUsesExponentialBackoff starts failing after a
// retry library bump, check whether the growth factor has changed.
func mulPhi(d time.Duration) time.Duration {
return time.Duration(float64(d) * math.Phi)
}
// setChatRunningAndPublish marks the chat row as running on workerID
// and publishes a matching status notification. It keeps the DB row
// and pubsub notification in sync so the async reconnect loop
// re-dials on each timer fire (the reconnect branch re-checks DB
// status before calling openRelayAsync).
func setChatRunningAndPublish(
ctx context.Context,
t *testing.T,
db database.Store,
ps dbpubsub.Pubsub,
chatID, workerID uuid.UUID,
) {
t.Helper()
now := time.Now()
_, err := db.UpdateChatStatus(ctx, database.UpdateChatStatusParams{
ID: chatID,
Status: database.ChatStatusRunning,
WorkerID: uuid.NullUUID{UUID: workerID, Valid: true},
StartedAt: sql.NullTime{Time: now, Valid: true},
HeartbeatAt: sql.NullTime{Time: now, Valid: true},
})
require.NoError(t, err)
payload, err := json.Marshal(coderdpubsub.ChatStreamNotifyMessage{
Status: string(database.ChatStatusRunning),
WorkerID: workerID.String(),
})
require.NoError(t, err)
require.NoError(t, ps.Publish(coderdpubsub.ChatStreamNotifyChannel(chatID), payload))
}
// TestRelayDialErrorIsUnrecoverable locks the classification policy.
// Adding a new HTTP status to the unrecoverable set should force a
// test edit too.
func TestRelayDialErrorIsUnrecoverable(t *testing.T) {
t.Parallel()
cases := []struct {
name string
status int
want bool
}{
{"unauthorized", http.StatusUnauthorized, true},
{"forbidden", http.StatusForbidden, true},
{"internal_server", http.StatusInternalServerError, false},
{"bad_gateway", http.StatusBadGateway, false},
{"service_unavailable", http.StatusServiceUnavailable, false},
{"too_many_requests", http.StatusTooManyRequests, false},
{"pre_response", 0, false},
{"bad_request", http.StatusBadRequest, false},
{"not_found", http.StatusNotFound, false},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
e := &entchatd.RelayDialError{HTTPStatus: tc.status, Err: io.EOF}
require.Equal(t, tc.want, e.IsUnrecoverable(),
"status=%d", tc.status)
})
}
}
// TestRelayReconnectUsesExponentialBackoff asserts that the reconnect
// timer follows the φ-growth sequence produced by
// github.com/coder/retry's defaults, floored at relayRetryFloor.
func TestRelayReconnectUsesExponentialBackoff(t *testing.T) {
t.Parallel()
db, ps := dbtestutil.NewDB(t)
workerID := uuid.New()
subscriberID := uuid.New()
var failCount atomic.Int32
dialer := func(_ context.Context, _ uuid.UUID, _ uuid.UUID, _ http.Header) (
[]codersdk.ChatStreamEvent, <-chan codersdk.ChatStreamEvent, func(), error,
) {
failCount.Add(1)
return nil, nil, nil, &entchatd.RelayDialError{
HTTPStatus: http.StatusBadGateway,
Err: io.EOF,
}
}
mclk := quartz.NewMock(t)
trapReconnect := mclk.Trap().NewTimer("reconnect")
defer trapReconnect.Close()
subscriber := newTestServer(t, db, ps, subscriberID, dialer, mclk)
ctx := testutil.Context(t, testutil.WaitLong)
user, org, model := seedChatDependencies(t, db)
chat := seedWaitingChat(t, db, org.ID, user, model, "relay-backoff")
_, events, cancel, ok := subscriber.Subscribe(ctx, chat.ID, nil, 0)
require.True(t, ok)
t.Cleanup(cancel)
// Kick the async relay loop and keep the DB row in sync so
// each reconnect timer fire triggers another dial.
setChatRunningAndPublish(ctx, t, db, ps, chat.ID, workerID)
// Expected sequence from retry.Retrier math:
// attempt 1 → floor (500ms)
// attempt n → prev × φ (capped at ceil)
floor := 500 * time.Millisecond
expected := []time.Duration{
floor,
mulPhi(floor),
mulPhi(mulPhi(floor)),
mulPhi(mulPhi(mulPhi(floor))),
mulPhi(mulPhi(mulPhi(mulPhi(floor)))),
}
for i, want := range expected {
call := trapReconnect.MustWait(ctx)
require.Equal(t, want, call.Duration,
"attempt %d: want %v got %v", i+1, want, call.Duration)
call.MustRelease(ctx)
mclk.Advance(want).MustWait(ctx)
}
// We expect 1 initial attempt + 5 reconnects fired by the
// trapped timer = 6 dials before the cap-check runs. Use
// Eventually so we don't race the final dial goroutine that
// the last Advance kicked off.
require.Eventually(t, func() bool {
return failCount.Load() >= 6
}, testutil.WaitShort, testutil.IntervalFast,
"expected 6 dials, got %d", failCount.Load())
// The events channel must remain open - we're still under the
// cap.
select {
case ev, open := <-events:
if !open {
t.Fatalf("events channel closed prematurely; retries should continue below cap")
}
// Allow through events that might have been queued; just
// confirm it's not a terminal error.
if ev.Type == codersdk.ChatStreamEventTypeError {
t.Fatalf("unexpected terminal error: %v", ev.Error)
}
default:
}
}
// TestRelayReconnectResetsOnSuccess exercises the path where a
// successful dial resets the retry state so the next failure starts
// over at the floor delay.
// TestRelayRepeatedDropsHitCap verifies the cap covers a peer that
// accepts the handshake and immediately drops it. Without a proper
// cap, such a peer would produce one reconnect per floor delay
// forever. The retry counter must accumulate across dial-success /
// parts-close cycles so the cap trips.
func TestRelayRepeatedDropsHitCap(t *testing.T) {
t.Parallel()
db, ps := dbtestutil.NewDB(t)
workerID := uuid.New()
subscriberID := uuid.New()
opened := make(chan chan codersdk.ChatStreamEvent, 32)
var call atomic.Int32
dialer := func(_ context.Context, _ uuid.UUID, _ uuid.UUID, _ http.Header) (
[]codersdk.ChatStreamEvent, <-chan codersdk.ChatStreamEvent, func(), error,
) {
call.Add(1)
ch := make(chan codersdk.ChatStreamEvent, 1)
opened <- ch
return nil, ch, func() {}, nil
}
mclk := quartz.NewMock(t)
trapReconnect := mclk.Trap().NewTimer("reconnect")
defer trapReconnect.Close()
subscriber := newTestServer(t, db, ps, subscriberID, dialer, mclk)
ctx := testutil.Context(t, testutil.WaitLong)
user, org, model := seedChatDependencies(t, db)
chat := seedWaitingChat(t, db, org.ID, user, model, "relay-drops")
_, events, cancel, ok := subscriber.Subscribe(ctx, chat.ID, nil, 0)
require.True(t, ok)
t.Cleanup(cancel)
// Kick off the first async dial.
setChatRunningAndPublish(ctx, t, db, ps, chat.ID, workerID)
// Close the first dial's parts channel so the merge loop
// schedules a reconnect. Then advance 6 reconnect timers,
// closing the parts channel each time so the cycle is:
// dial -> success -> parts-close -> next() -> reconnect.
// 1 initial dial + 6 timer-driven dials = 7 total; the 7th
// parts-close trips the cap.
for i := 0; i < 7; i++ {
var ch chan codersdk.ChatStreamEvent
select {
case ch = <-opened:
case <-ctx.Done():
t.Fatalf("timed out waiting for dial %d", i+1)
}
// Closing the parts channel triggers the relayPartsCh
// close branch, which calls retryState.next() and
// schedules the next reconnect.
close(ch)
if i == 6 {
// 7th parts-close should trip the cap; no more
// reconnect timers.
break
}
call := trapReconnect.MustWait(ctx)
call.MustRelease(ctx)
mclk.Advance(call.Duration).MustWait(ctx)
}
// A terminal error event must arrive on the events channel.
var errEvent *codersdk.ChatStreamEvent
require.Eventually(t, func() bool {
select {
case ev, open := <-events:
if !open {
return errEvent != nil
}
if ev.Type == codersdk.ChatStreamEventTypeError {
errEvent = &ev
return true
}
return false
default:
return false
}
}, testutil.WaitShort, testutil.IntervalFast,
"expected a terminal error event after repeated drops hit cap")
require.NotNil(t, errEvent.Error)
require.Contains(t, errEvent.Error.Message, "relay connection failed")
// We should have observed exactly 7 dials before tear-down.
require.Equal(t, int32(7), call.Load(),
"expected 7 dials (1 initial + 6 reconnect retries) before cap")
}
// TestRelayStopsAfterIntermittentCap verifies the cap-reached
// tear-down path: after N intermittent failures the merge loop emits
// one error event, closes the events channel, and stops dialing.
func TestRelayStopsAfterIntermittentCap(t *testing.T) {
t.Parallel()
db, ps := dbtestutil.NewDB(t)
workerID := uuid.New()
subscriberID := uuid.New()
var callCount atomic.Int32
dialer := func(_ context.Context, _ uuid.UUID, _ uuid.UUID, _ http.Header) (
[]codersdk.ChatStreamEvent, <-chan codersdk.ChatStreamEvent, func(), error,
) {
callCount.Add(1)
return nil, nil, nil, &entchatd.RelayDialError{
HTTPStatus: http.StatusBadGateway,
Err: io.EOF,
}
}
mclk := quartz.NewMock(t)
trapReconnect := mclk.Trap().NewTimer("reconnect")
defer trapReconnect.Close()
subscriber := newTestServer(t, db, ps, subscriberID, dialer, mclk)
ctx := testutil.Context(t, testutil.WaitLong)
user, org, model := seedChatDependencies(t, db)
chat := seedWaitingChat(t, db, org.ID, user, model, "relay-cap")
_, events, cancel, ok := subscriber.Subscribe(ctx, chat.ID, nil, 0)
require.True(t, ok)
t.Cleanup(cancel)
setChatRunningAndPublish(ctx, t, db, ps, chat.ID, workerID)
// Advance through N consecutive reconnect timers. Each one
// triggers a dial, which fails and schedules the next timer.
// After the Nth failure the retry state says giveUp=true on
// the next .next() call, so the merge loop tears down.
for i := 0; i < 6; i++ {
call := trapReconnect.MustWait(ctx)
call.MustRelease(ctx)
mclk.Advance(call.Duration).MustWait(ctx)
}
// Wait for the terminal error event to arrive. mergedEvents
// closes inside the enterprise merge goroutine, but OSS only
// nil-outs relayEvents on close - the outer events channel
// stays open for pubsub/local, so we wait for the error event
// itself rather than channel closure.
var errEvent *codersdk.ChatStreamEvent
require.Eventually(t, func() bool {
select {
case ev, open := <-events:
if !open {
return errEvent != nil
}
if ev.Type == codersdk.ChatStreamEventTypeError {
errEvent = &ev
return true
}
return false
default:
return false
}
}, testutil.WaitShort, testutil.IntervalFast,
"expected a terminal error event")
require.NotNil(t, errEvent, "expected a terminal error event")
require.NotNil(t, errEvent.Error)
require.Contains(t, errEvent.Error.Message, "relay connection failed")
require.Contains(t, errEvent.Error.Message, "6")
// Ensure the cap fires at attempt N+1 - the retry state allows
// relayMaxRetries successful next() calls before flipping
// giveUp. With one initial dial + 6 reconnect-timer fires the
// 7th .next() trips the cap and tears down, so we see 7 dials
// total and nothing further.
totalDials := callCount.Load()
require.Equal(t, int32(7), totalDials,
"expected exactly relayMaxRetries+1 dials before cap; got %d", totalDials)
}
// chatByIDErrorStore wraps a database.Store and forces GetChatByID
// to return a caller-supplied error once after N successful calls.
// This lets the initial Subscribe call succeed (OSS's initial state
// load needs a real Chat to wire up the relay) while subsequent
// reconnect-branch calls exercise the DB-error retry path.
type chatByIDErrorStore struct {
database.Store
err error
okRemain atomic.Int32 // number of calls allowed to delegate before erroring.
}
func (s *chatByIDErrorStore) GetChatByID(ctx context.Context, id uuid.UUID) (database.Chat, error) {
if s.okRemain.Add(-1) >= 0 {
return s.Store.GetChatByID(ctx, id)
}
return database.Chat{}, s.err
}
// TestRelayReconnectStopsAfterDBErrorCap verifies the reconnect-timer
// branch's DB-error path shares the same retry budget as dial
// failures and trips the cap after enough consecutive DB errors.
func TestRelayReconnectStopsAfterDBErrorCap(t *testing.T) {
t.Parallel()
realDB, ps := dbtestutil.NewDB(t)
workerID := uuid.New()
subscriberID := uuid.New()
var callCount atomic.Int32
dialer := func(_ context.Context, _ uuid.UUID, _ uuid.UUID, _ http.Header) (
[]codersdk.ChatStreamEvent, <-chan codersdk.ChatStreamEvent, func(), error,
) {
callCount.Add(1)
return nil, nil, nil, &entchatd.RelayDialError{
HTTPStatus: http.StatusBadGateway,
Err: io.EOF,
}
}
mclk := quartz.NewMock(t)
trapReconnect := mclk.Trap().NewTimer("reconnect")
defer trapReconnect.Close()
// The server sees a DB whose GetChatByID always errors after
// the initial Subscribe snapshot load. Other methods delegate
// to the real DB, so seeding below still works.
failingDB := &chatByIDErrorStore{
Store: realDB,
err: xerrors.New("mock: GetChatByID always fails"),
}
// Allow one successful GetChatByID (the Subscribe preamble's
// initial state load). All subsequent calls return the mock
// error, exercising the reconnect-branch DB-error path.
failingDB.okRemain.Store(1)
ctx := testutil.Context(t, testutil.WaitLong)
user, org, model := seedChatDependencies(t, realDB)
chat := seedWaitingChat(t, realDB, org.ID, user, model, "relay-db-error")
subscriber := newTestServer(t, failingDB, ps, subscriberID, dialer, mclk)
_, events, cancel, ok := subscriber.Subscribe(ctx, chat.ID, nil, 0)
require.True(t, ok)
t.Cleanup(cancel)
// Flip to running so the merge loop starts an async dial. The
// dial fails (attempts=1, reconnect scheduled). From there each
// reconnect timer fires, the merge loop calls GetChatByID, the
// failing DB returns an error, and retryState.next() increments.
//
// Budget: 1 dial-failure + 6 DB-failures = 7 next() calls; the
// 7th trips the cap.
setChatRunningAndPublish(ctx, t, realDB, ps, chat.ID, workerID)
for i := 0; i < 6; i++ {
call := trapReconnect.MustWait(ctx)
call.MustRelease(ctx)
mclk.Advance(call.Duration).MustWait(ctx)
}
var errEvent *codersdk.ChatStreamEvent
require.Eventually(t, func() bool {
select {
case ev, open := <-events:
if !open {
return errEvent != nil
}
if ev.Type == codersdk.ChatStreamEventTypeError {
errEvent = &ev
return true
}
return false
default:
return false
}
}, testutil.WaitShort, testutil.IntervalFast,
"expected terminal error event after DB-error cap")
require.NotNil(t, errEvent.Error)
require.Contains(t, errEvent.Error.Message, "relay connection failed")
require.Contains(t, errEvent.Error.Message, "6")
// Exactly 1 dial fired: the one that triggered the initial
// reconnect schedule. All subsequent next() calls come from the
// DB-error branch without calling the dialer.
require.Equal(t, int32(1), callCount.Load(),
"expected exactly 1 dial; reconnects should short-circuit on DB error")
}
// TestRelayStopsImmediatelyOnUnauthorized tests the unrecoverable
// branch and its table of status codes.
func TestRelayStopsImmediatelyOnUnauthorized(t *testing.T) {
t.Parallel()
cases := []struct {
name string
status int
wantUnrecoverable bool
wantMsgContains string
}{
{"401", http.StatusUnauthorized, true, "401"},
{"403", http.StatusForbidden, true, "403"},
{"500_intermittent", http.StatusInternalServerError, false, ""},
{"zero_intermittent", 0, false, ""},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
db, ps := dbtestutil.NewDB(t)
workerID := uuid.New()
subscriberID := uuid.New()
var callCount atomic.Int32
dialer := func(_ context.Context, _ uuid.UUID, _ uuid.UUID, _ http.Header) (
[]codersdk.ChatStreamEvent, <-chan codersdk.ChatStreamEvent, func(), error,
) {
callCount.Add(1)
return nil, nil, nil, &entchatd.RelayDialError{
HTTPStatus: tc.status,
Err: io.EOF,
}
}
mclk := quartz.NewMock(t)
trapReconnect := mclk.Trap().NewTimer("reconnect")
defer trapReconnect.Close()
subscriber := newTestServer(t, db, ps, subscriberID, dialer, mclk)
ctx := testutil.Context(t, testutil.WaitLong)
user, org, model := seedChatDependencies(t, db)
chat := seedWaitingChat(t, db, org.ID, user, model,
"relay-unrec-"+tc.name)
_, events, cancel, ok := subscriber.Subscribe(ctx, chat.ID, nil, 0)
require.True(t, ok)
t.Cleanup(cancel)
setChatRunningAndPublish(ctx, t, db, ps, chat.ID, workerID)
if tc.wantUnrecoverable {
// First dial should tear the relay down.
var errEvent *codersdk.ChatStreamEvent
require.Eventually(t, func() bool {
select {
case ev, open := <-events:
if !open {
return errEvent != nil
}
if ev.Type == codersdk.ChatStreamEventTypeError {
errEvent = &ev
return true
}
return false
default:
return false
}
}, testutil.WaitShort, testutil.IntervalFast,
"expected terminal error event")
require.NotNil(t, errEvent)
require.Contains(t, errEvent.Error.Message, "relay authentication failed")
require.Contains(t, errEvent.Error.Message, tc.wantMsgContains)
require.Equal(t, int32(1), callCount.Load(),
"unrecoverable errors must not retry; got %d dials", callCount.Load())
} else {
// Intermittent: fire one reconnect timer
// and confirm the dialer is called again.
call := trapReconnect.MustWait(ctx)
call.MustRelease(ctx)
mclk.Advance(call.Duration).MustWait(ctx)
require.Eventually(t, func() bool {
return callCount.Load() >= 2
}, testutil.WaitShort, testutil.IntervalFast,
"intermittent should retry at least once")
}
})
}
}
// TestRelayBackoffResetsOnStatusChange checks that closeRelay (driven
// by a status notification) resets the retry counter so subsequent
// dials against a new target start at the floor delay.
func TestRelayBackoffResetsOnStatusChange(t *testing.T) {
t.Parallel()
db, ps := dbtestutil.NewDB(t)
workerID1 := uuid.New()
workerID2 := uuid.New()
subscriberID := uuid.New()
dialer := func(_ context.Context, _ uuid.UUID, _ uuid.UUID, _ http.Header) (
[]codersdk.ChatStreamEvent, <-chan codersdk.ChatStreamEvent, func(), error,
) {
return nil, nil, nil, &entchatd.RelayDialError{
HTTPStatus: http.StatusBadGateway,
Err: io.EOF,
}
}
mclk := quartz.NewMock(t)
trapReconnect := mclk.Trap().NewTimer("reconnect")
defer trapReconnect.Close()
subscriber := newTestServer(t, db, ps, subscriberID, dialer, mclk)
ctx := testutil.Context(t, testutil.WaitLong)
user, org, model := seedChatDependencies(t, db)
chat := seedWaitingChat(t, db, org.ID, user, model, "relay-reset-on-status")
_, _, cancel, ok := subscriber.Subscribe(ctx, chat.ID, nil, 0)
require.True(t, ok)
t.Cleanup(cancel)
// Drive the async openRelayAsync path with workerID1.
setChatRunningAndPublish(ctx, t, db, ps, chat.ID, workerID1)
// Drive 3 intermittent failures so attempts=3 and the delay
// has grown past the floor. After each loop iteration the 4th
// reconnect timer is queued - consume it too so our later
// assertion sees the reset's timer, not a stale one.
for i := 0; i < 3; i++ {
call := trapReconnect.MustWait(ctx)
call.MustRelease(ctx)
mclk.Advance(call.Duration).MustWait(ctx)
}
// Grab the next trapped timer (the grown one scheduled after
// the 3rd dial fails) but don't advance it - we want to see it
// replaced by a fresh floor-delay timer after the reset.
grown := trapReconnect.MustWait(ctx)
require.Greater(t, grown.Duration, 500*time.Millisecond,
"sanity: pre-reset delay should have grown past the floor")
grown.MustRelease(ctx)
// Flip the chat to waiting; closeRelay runs (because the
// status notification no longer points at a running peer) and
// should reset the retry state.
_, err := db.UpdateChatStatus(ctx, database.UpdateChatStatusParams{
ID: chat.ID,
Status: database.ChatStatusWaiting,
})
require.NoError(t, err)
waitingPayload, err := json.Marshal(coderdpubsub.ChatStreamNotifyMessage{
Status: string(database.ChatStatusWaiting),
})
require.NoError(t, err)
require.NoError(t, ps.Publish(coderdpubsub.ChatStreamNotifyChannel(chat.ID), waitingPayload))
// Flip back to running on a different worker. This triggers a
// fresh openRelayAsync which fails, arming a reconnect timer.
// That timer's delay must be the floor, proving the reset.
setChatRunningAndPublish(ctx, t, db, ps, chat.ID, workerID2)
call := trapReconnect.MustWait(ctx)
require.Equal(t, 500*time.Millisecond, call.Duration,
"retry state must reset after status change; got grown delay %v", call.Duration)
call.MustRelease(ctx)
}
// TestRelayBackoffRespectsContextCancel is a regression guard: the
// reconnect timer must respect ctx cancellation promptly.
func TestRelayBackoffRespectsContextCancel(t *testing.T) {
t.Parallel()
db, ps := dbtestutil.NewDB(t)
workerID := uuid.New()
subscriberID := uuid.New()
dialer := func(_ context.Context, _ uuid.UUID, _ uuid.UUID, _ http.Header) (
[]codersdk.ChatStreamEvent, <-chan codersdk.ChatStreamEvent, func(), error,
) {
return nil, nil, nil, &entchatd.RelayDialError{
HTTPStatus: http.StatusBadGateway,
Err: io.EOF,
}
}
mclk := quartz.NewMock(t)
trapReconnect := mclk.Trap().NewTimer("reconnect")
defer trapReconnect.Close()
subscriber := newTestServer(t, db, ps, subscriberID, dialer, mclk)
ctx := testutil.Context(t, testutil.WaitLong)
user, org, model := seedChatDependencies(t, db)
chat := seedWaitingChat(t, db, org.ID, user, model, "relay-cancel")
subCtx, subCancel := context.WithCancel(ctx)
_, events, cancel, ok := subscriber.Subscribe(subCtx, chat.ID, nil, 0)
require.True(t, ok)
t.Cleanup(cancel)
setChatRunningAndPublish(ctx, t, db, ps, chat.ID, workerID)
// Wait for the first reconnect timer to arm.
call := trapReconnect.MustWait(ctx)
call.MustRelease(ctx)
// Cancel the subscriber context. The events channel should
// close promptly (the merge goroutine's select exits on
// ctx.Done).
subCancel()
done := make(chan struct{})
go func() {
defer close(done)
for {
if _, open := <-events; !open {
return
}
}
}()
select {
case <-done:
case <-time.After(testutil.WaitShort):
t.Fatal("events channel did not close after ctx cancel")
}
}
// TestDialRelayReal401 exercises the real dialRelay path against an
// httptest server that returns 401 on the stream endpoint. It
// validates that the websocket library's handshake failure
// propagates through as *RelayDialError with HTTPStatus == 401.
//
// This is the one test that uses the real coder/websocket library
// on the failure path - a safety net against library upgrades
// silently breaking status capture.
func TestDialRelayReal401(t *testing.T) {
t.Parallel()
// An httptest server that 401s every request on the stream
// endpoint. Any other path gets a 404.
srv := httptest.NewServer(http.HandlerFunc(func(rw http.ResponseWriter, r *http.Request) {
if !streamPathRE.MatchString(r.URL.Path) {
http.NotFound(rw, r)
return
}
rw.Header().Set("Content-Type", "application/json")
rw.WriteHeader(http.StatusUnauthorized)
_, _ = rw.Write([]byte(`{"message":"unauthorized"}`))
}))
t.Cleanup(srv.Close)
db, _ := dbtestutil.NewDB(t)
workerID := uuid.New()
subscriberID := uuid.New()
// Wire real config (no DialerFn override) so dialRelay runs
// end-to-end against the httptest server. Seeding a waiting
// chat (below) keeps Subscribe's initial synchronous dial a
// no-op; we then push a running status notification to the
// merge loop so it invokes dialRelay via the async path, where
// the 401 tear-down logic lives.
cfg := entchatd.MultiReplicaSubscribeConfig{
ResolveReplicaAddress: func(_ context.Context, _ uuid.UUID) (string, bool) {
return srv.URL, true
},
ReplicaHTTPClient: srv.Client(),
ReplicaIDFn: func() uuid.UUID { return subscriberID },
}
subscribeFn := entchatd.NewMultiReplicaSubscribeFn(cfg)
ctx := testutil.Context(t, testutil.WaitMedium)
user, org, model := seedChatDependencies(t, db)
// Seed a waiting chat - no sync dial - then push a running
// status notification to trigger the async dial via the real
// dialRelay path.
chat := seedWaitingChat(t, db, org.ID, user, model, "relay-real-401")
statusCh := make(chan osschatd.StatusNotification, 1)
evs := subscribeFn(ctx, osschatd.SubscribeFnParams{
ChatID: chat.ID,
Chat: chat,
WorkerID: subscriberID,
StatusNotifications: statusCh,
RequestHeader: http.Header{codersdk.SessionTokenHeader: {"test-token"}},
DB: db,
Logger: slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}),
})
statusCh <- osschatd.StatusNotification{
Status: database.ChatStatusRunning,
WorkerID: workerID,
}
// Wait for a terminal error event. On a real 401 handshake,
// the classifier flags it unrecoverable → one dial, then
// error event, then channel close.
var errEvent *codersdk.ChatStreamEvent
deadline := time.After(testutil.WaitMedium)
waitErr:
for {
select {
case ev, open := <-evs:
if !open {
break waitErr
}
if ev.Type == codersdk.ChatStreamEventTypeError {
errEvent = &ev
}
case <-deadline:
break waitErr
}
}
require.NotNil(t, errEvent, "expected terminal error event from real 401 dial")
require.NotNil(t, errEvent.Error)
require.Contains(t, errEvent.Error.Message, "relay authentication failed")
require.Contains(t, errEvent.Error.Message, "401")
}
// streamPathRE matches the chat stream endpoint path built by
// buildRelayURL. Compiled at package scope so the httptest handler
// below doesn't pay regexp.Compile per request.
var streamPathRE = regexp.MustCompile(
`^/api/experimental/chats/[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}/stream$`,
)
File diff suppressed because it is too large Load Diff