diff --git a/cli/server.go b/cli/server.go index 0df991e56e..9a36ff6e9b 100644 --- a/cli/server.go +++ b/cli/server.go @@ -2376,6 +2376,19 @@ func redirectToAccessURL(handler http.Handler, accessURL *url.URL, tunnel bool, return } + // Exception: inter-replica relay. + // Enterprise chat streaming relays message_part events + // between replicas by dialing the worker replica's + // DERP relay address directly. Redirecting these + // requests to the access URL breaks the WebSocket + // handshake because the redirect strips the Upgrade + // headers, causing the load-balanced access URL to + // return HTTP 200 (SPA catch-all) instead of 101. + if isReplicaRelayRequest(r) { + handler.ServeHTTP(w, r) + return + } + // Only do this if we aren't tunneling. // If we are tunneling, we want to allow the request to go through // because the tunnel doesn't proxy with TLS. @@ -2411,6 +2424,14 @@ func isDERPPath(p string) bool { return segments[1] == "derp" } +// isReplicaRelayRequest returns true when the request was sent by +// another coderd replica as part of cross-replica streaming. The +// enterprise chat relay sets X-Coder-Relay-Source-Replica on every +// request to identify itself. +func isReplicaRelayRequest(r *http.Request) bool { + return r.Header.Get("X-Coder-Relay-Source-Replica") != "" +} + // IsLocalhost returns true if the host points to the local machine. Intended to // be called with `u.Hostname()`. func IsLocalhost(host string) bool { diff --git a/cli/server_internal_test.go b/cli/server_internal_test.go index 22a53d030b..e2f5b8df32 100644 --- a/cli/server_internal_test.go +++ b/cli/server_internal_test.go @@ -4,6 +4,7 @@ import ( "bytes" "context" "crypto/tls" + "net/http" "testing" "github.com/spf13/pflag" @@ -314,6 +315,30 @@ func TestIsDERPPath(t *testing.T) { } } +func TestIsReplicaRelayRequest(t *testing.T) { + t.Parallel() + + t.Run("WithHeader", func(t *testing.T) { + t.Parallel() + r, _ := http.NewRequestWithContext(context.Background(), "GET", "/api/experimental/chats/abc/stream", nil) + r.Header.Set("X-Coder-Relay-Source-Replica", "some-uuid") + require.True(t, isReplicaRelayRequest(r)) + }) + + t.Run("WithoutHeader", func(t *testing.T) { + t.Parallel() + r, _ := http.NewRequestWithContext(context.Background(), "GET", "/api/experimental/chats/abc/stream", nil) + require.False(t, isReplicaRelayRequest(r)) + }) + + t.Run("EmptyHeader", func(t *testing.T) { + t.Parallel() + r, _ := http.NewRequestWithContext(context.Background(), "GET", "/api/experimental/chats/abc/stream", nil) + r.Header.Set("X-Coder-Relay-Source-Replica", "") + require.False(t, isReplicaRelayRequest(r)) + }) +} + func TestEscapePostgresURLUserInfo(t *testing.T) { t.Parallel()