Thanks to visit codestin.com
Credit goes to github.com

Skip to content

Commit d8e69d3

Browse files
authored
fix: enable Copilot HTTP transport fallback (#28494) (#28722)
Backport of #28494 Original PR: #28494 — fix: enable Copilot HTTP transport fallback Merge commit: 849543d Requested by: @ssncferreira > [!NOTE] > This pull request was generated by Coder Agents on behalf of @ssncferreira.
1 parent 5a07bdf commit d8e69d3

8 files changed

Lines changed: 335 additions & 70 deletions

File tree

aibridge/bridge.go

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -18,6 +18,7 @@ import (
1818
"github.com/sony/gobreaker/v2"
1919
"go.opentelemetry.io/otel/codes"
2020
"go.opentelemetry.io/otel/trace"
21+
"golang.org/x/net/http/httpguts"
2122
"golang.org/x/xerrors"
2223

2324
"cdr.dev/slog/v3"
@@ -248,6 +249,18 @@ func newInterceptionProcessor(p provider.Provider, cbs *circuitbreaker.ProviderC
248249
client := GuessClient(r)
249250
sessionID := GuessSessionID(client, r)
250251

252+
if isWebSocketUpgrade(r) {
253+
route := strings.TrimPrefix(r.URL.Path, fmt.Sprintf("/%s", p.Name()))
254+
logger.Debug(ctx, "rejecting unsupported WebSocket upgrade",
255+
slog.F("provider", p.Name()),
256+
slog.F("route", route),
257+
slog.F("client", string(client)),
258+
slog.F("client_session_id", sessionID),
259+
)
260+
http.Error(w, "WebSocket transport is not supported, use HTTP", http.StatusNotImplemented)
261+
return
262+
}
263+
251264
// Read and validate Agent Firewall correlation headers. The
252265
// values are captured here and recorded below; the headers
253266
// themselves are stripped from the upstream request by
@@ -378,6 +391,13 @@ func newInterceptionProcessor(p provider.Provider, cbs *circuitbreaker.ProviderC
378391
}
379392
}
380393

394+
// isWebSocketUpgrade reports whether r is a WebSocket opening handshake.
395+
func isWebSocketUpgrade(r *http.Request) bool {
396+
return r.Method == http.MethodGet &&
397+
httpguts.HeaderValuesContainsToken(r.Header.Values("Connection"), "upgrade") &&
398+
httpguts.HeaderValuesContainsToken(r.Header.Values("Upgrade"), "websocket")
399+
}
400+
381401
// writeRequestBodyTooLarge writes a human-readable 413 response indicating that
382402
// the request body exceeded maxRequestBodyBytes.
383403
func writeRequestBodyTooLarge(w http.ResponseWriter) {

aibridge/bridge_internal_test.go

Lines changed: 30 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,36 @@ import (
1010
agplaibridge "github.com/coder/coder/v2/coderd/aibridge"
1111
)
1212

13+
func TestIsWebSocketUpgrade(t *testing.T) {
14+
t.Parallel()
15+
16+
tests := []struct {
17+
name string
18+
method string
19+
connection string
20+
upgrade string
21+
want bool
22+
}{
23+
{name: "websocket upgrade", method: http.MethodGet, connection: "keep-alive, Upgrade", upgrade: "WebSocket", want: true},
24+
{name: "non-GET request", method: http.MethodPost, connection: "Upgrade", upgrade: "websocket", want: false},
25+
{name: "missing connection upgrade", method: http.MethodGet, connection: "keep-alive", upgrade: "websocket", want: false},
26+
{name: "different upgrade protocol", method: http.MethodGet, connection: "Upgrade", upgrade: "h2c", want: false},
27+
}
28+
29+
for _, tc := range tests {
30+
t.Run(tc.name, func(t *testing.T) {
31+
t.Parallel()
32+
33+
req, err := http.NewRequestWithContext(t.Context(), tc.method, "/", nil)
34+
require.NoError(t, err)
35+
req.Header.Set("Connection", tc.connection)
36+
req.Header.Set("Upgrade", tc.upgrade)
37+
38+
assert.Equal(t, tc.want, isWebSocketUpgrade(req))
39+
})
40+
}
41+
}
42+
1343
func TestExtractAgentFirewallHeaders(t *testing.T) {
1444
t.Parallel()
1545

aibridge/bridge_test.go

Lines changed: 57 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -13,11 +13,13 @@ import (
1313
"github.com/stretchr/testify/assert"
1414
"github.com/stretchr/testify/require"
1515
"go.opentelemetry.io/otel"
16+
"go.opentelemetry.io/otel/trace"
1617

1718
"cdr.dev/slog/v3/sloggers/slogtest"
1819
"github.com/coder/coder/v2/aibridge"
1920
"github.com/coder/coder/v2/aibridge/aibridgetest"
2021
"github.com/coder/coder/v2/aibridge/config"
22+
"github.com/coder/coder/v2/aibridge/intercept"
2123
"github.com/coder/coder/v2/aibridge/internal/testutil"
2224
"github.com/coder/coder/v2/aibridge/provider"
2325
codertestutil "github.com/coder/coder/v2/testutil"
@@ -186,11 +188,12 @@ func TestPassthroughRoutesForProviders(t *testing.T) {
186188

187189
upstreamRespBody := "upstream response"
188190
tests := []struct {
189-
name string
190-
baseURLPath string
191-
requestPath string
192-
provider func(*testing.T, string) provider.Provider
193-
expectPath string
191+
name string
192+
baseURLPath string
193+
requestMethod string
194+
requestPath string
195+
provider func(*testing.T, string) provider.Provider
196+
expectPath string
194197
}{
195198
{
196199
name: "openAI_no_base_path",
@@ -243,6 +246,23 @@ func TestPassthroughRoutesForProviders(t *testing.T) {
243246
},
244247
expectPath: "/v1/models",
245248
},
249+
{
250+
name: "copilot_ping",
251+
requestPath: "/copilot/_ping",
252+
provider: func(_ *testing.T, baseURL string) provider.Provider {
253+
return aibridge.NewCopilotProvider(config.Copilot{BaseURL: baseURL})
254+
},
255+
expectPath: "/_ping",
256+
},
257+
{
258+
name: "copilot_auto",
259+
requestMethod: http.MethodPost,
260+
requestPath: "/copilot/auto",
261+
provider: func(_ *testing.T, baseURL string) provider.Provider {
262+
return aibridge.NewCopilotProvider(config.Copilot{BaseURL: baseURL})
263+
},
264+
expectPath: "/auto",
265+
},
246266
}
247267

248268
for _, tc := range tests {
@@ -263,7 +283,7 @@ func TestPassthroughRoutesForProviders(t *testing.T) {
263283
bridge, err := aibridge.NewRequestBridge(t.Context(), []provider.Provider{prov}, &rec, nil, logger, nil, bridgeTestTracer)
264284
require.NoError(t, err)
265285

266-
req := httptest.NewRequest("", tc.requestPath, nil)
286+
req := httptest.NewRequest(tc.requestMethod, tc.requestPath, nil)
267287
resp := httptest.NewRecorder()
268288
bridge.ServeHTTP(resp, req)
269289

@@ -273,6 +293,37 @@ func TestPassthroughRoutesForProviders(t *testing.T) {
273293
}
274294
}
275295

296+
func TestWebSocketUpgradeRejected(t *testing.T) {
297+
t.Parallel()
298+
299+
interceptorCalled := false
300+
prov := &testutil.MockProvider{
301+
NameStr: "test",
302+
Bridged: []string{"/responses"},
303+
InterceptorFunc: func(http.ResponseWriter, *http.Request, trace.Tracer) (intercept.Interceptor, error) {
304+
interceptorCalled = true
305+
return nil, nil //nolint:nilnil // The interceptor must not be reached.
306+
},
307+
}
308+
bridge, err := aibridge.NewRequestBridge(
309+
t.Context(),
310+
[]provider.Provider{prov},
311+
nil, nil, slogtest.Make(t, nil), nil, bridgeTestTracer,
312+
)
313+
require.NoError(t, err)
314+
315+
req := httptest.NewRequest(http.MethodGet, "/test/responses", nil)
316+
req.Header.Set("Connection", "keep-alive, Upgrade")
317+
req.Header.Set("Upgrade", "WebSocket")
318+
resp := httptest.NewRecorder()
319+
320+
bridge.ServeHTTP(resp, req)
321+
322+
assert.Equal(t, http.StatusNotImplemented, resp.Code)
323+
assert.Contains(t, resp.Body.String(), "WebSocket transport is not supported, use HTTP")
324+
assert.False(t, interceptorCalled)
325+
}
326+
276327
func TestRequestBodySizeLimit(t *testing.T) {
277328
t.Parallel()
278329

aibridge/provider/copilot.go

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -87,6 +87,8 @@ func (*Copilot) BridgedRoutes() []string {
8787

8888
func (*Copilot) PassthroughRoutes() []string {
8989
return []string{
90+
"/_ping",
91+
"/auto",
9092
"/models",
9193
"/models/",
9294
"/agents/",

enterprise/aibridgeproxyd/aibridgeproxyd.go

Lines changed: 45 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -28,6 +28,7 @@ import (
2828
"golang.org/x/xerrors"
2929

3030
"cdr.dev/slog/v3"
31+
aibridgeconfig "github.com/coder/coder/v2/aibridge/config"
3132
agplaibridge "github.com/coder/coder/v2/coderd/aibridge"
3233
)
3334

@@ -132,7 +133,7 @@ type Server struct {
132133
// refreshProviders fetches the live provider snapshot on Reload.
133134
// Nil disables hot-reload.
134135
refreshProviders RefreshProvidersFunc
135-
// providerRouter holds the live (mitmHosts, nameByHost) pair.
136+
// providerRouter holds the live routing snapshot.
136137
providerRouter atomic.Pointer[providerRouter]
137138
// allowedPorts is the port allowlist for CONNECT requests. Fixed at
138139
// construction; not reloadable.
@@ -149,19 +150,26 @@ type Server struct {
149150
metrics *Metrics
150151
}
151152

153+
type routedProvider struct {
154+
name string
155+
providerType string
156+
}
157+
152158
// providerRouter keeps CONNECT matching and provider lookup in sync.
153159
type providerRouter struct {
154-
mitmHosts []string // host:port set the goproxy condition matches against.
155-
nameByHost map[string]string // lowercase hostname -> provider name.
160+
mitmHosts []string // host:port set the goproxy condition matches against.
161+
providerByHost map[string]routedProvider // lowercase hostname -> provider.
156162
}
157163

158164
// emptyProviderRouter is used before the first Reload (or when the
159165
// operator deconfigures every provider) so handlers can safely call
160166
// loadProviderRouter without a nil check.
161-
var emptyProviderRouter = &providerRouter{nameByHost: map[string]string{}}
167+
var emptyProviderRouter = &providerRouter{
168+
providerByHost: map[string]routedProvider{},
169+
}
162170

163-
func (r *providerRouter) providerFromHost(host string) string {
164-
return r.nameByHost[strings.ToLower(host)]
171+
func (r *providerRouter) providerFromHost(host string) routedProvider {
172+
return r.providerByHost[strings.ToLower(host)]
165173
}
166174

167175
// requestContext holds metadata propagated through the proxy request/response chain.
@@ -651,13 +659,13 @@ func (s *Server) authMiddleware(host string, ctx *goproxy.ProxyCtx) (*goproxy.Co
651659
provider := s.loadProviderRouter().providerFromHost(ctx.Req.URL.Hostname())
652660
// A concurrent Reload can swap the router between CONNECT matching
653661
// and provider lookup, so treat a missing mapping as a runtime miss.
654-
if provider == "" {
662+
if provider.name == "" {
655663
logger.Warn(s.ctx, "rejecting CONNECT request with no provider mapping")
656664
return goproxy.RejectConnect, host
657665
}
658666

659667
logger = logger.With(
660-
slog.F("provider", provider),
668+
slog.F("provider", provider.name),
661669
)
662670

663671
proxyAuth := ctx.Req.Header.Get("Proxy-Authorization")
@@ -681,7 +689,7 @@ func (s *Server) authMiddleware(host string, ctx *goproxy.ProxyCtx) (*goproxy.Co
681689
ctx.UserData = &requestContext{
682690
ConnectSessionID: connectSessionID,
683691
CoderToken: coderToken,
684-
Provider: provider,
692+
Provider: provider.name,
685693
}
686694

687695
logger.Debug(s.ctx, "request CONNECT authenticated")
@@ -932,14 +940,14 @@ func (s *Server) handleRequest(req *http.Request, ctx *goproxy.ProxyCtx) (*http.
932940
}
933941
}
934942
liveProvider := s.loadProviderRouter().providerFromHost(host)
935-
if liveProvider == "" || liveProvider != reqCtx.Provider {
943+
if liveProvider.name == "" || liveProvider.name != reqCtx.Provider {
936944
s.logger.Warn(s.ctx, "provider mapping changed or removed since CONNECT, passing through",
937945
slog.F("connect_id", reqCtx.ConnectSessionID.String()),
938946
slog.F("host", req.Host),
939947
slog.F("method", req.Method),
940948
slog.F("path", originalPath),
941949
slog.F("connect_provider", reqCtx.Provider),
942-
slog.F("live_provider", liveProvider),
950+
slog.F("live_provider", liveProvider.name),
943951
)
944952
return req, nil
945953
}
@@ -988,7 +996,8 @@ func (s *Server) handleRequest(req *http.Request, ctx *goproxy.ProxyCtx) (*http.
988996
req.URL = aiBridgeParsedURL
989997
req.Host = aiBridgeParsedURL.Host
990998

991-
injectBYOKHeaderIfNeeded(req.Header, reqCtx.CoderToken)
999+
// Prepare Coder authentication for centralized and BYOK requests.
1000+
prepareAIGatewayAuth(req.Header, reqCtx.CoderToken, liveProvider.providerType)
9921001

9931002
// Set request ID header to correlate requests between aibridgeproxyd and aibridged.
9941003
req.Header.Set(agplaibridge.HeaderCoderRequestID, reqCtx.RequestID.String())
@@ -1015,24 +1024,34 @@ func (s *Server) handleRequest(req *http.Request, ctx *goproxy.ProxyCtx) (*http.
10151024
return req, nil
10161025
}
10171026

1018-
// injectBYOKHeaderIfNeeded sets HeaderCoderToken when the
1019-
// Authorization header carries a bearer token that differs from the
1020-
// Coder token, indicating the client is using its own LLM
1021-
// credentials. Clients that can set custom headers
1022-
// do this themselves; this handles clients that cannot.
1023-
//
1024-
// In centralized mode, Authorization carries the Coder token
1025-
// itself, so aibridged discovers it via ExtractAuthToken
1026-
// without any extra header.
1027-
func injectBYOKHeaderIfNeeded(header http.Header, coderToken string) {
1028-
// Don’t overwrite the header if it’s already set.
1029-
if header.Get(agplaibridge.HeaderCoderToken) != "" {
1027+
// prepareAIGatewayAuth prepares the Coder authentication headers for AI
1028+
// Gateway. Copilot is always BYOK, while other providers may use centralized
1029+
// or BYOK authentication.
1030+
func prepareAIGatewayAuth(headers http.Header, coderToken, providerType string) {
1031+
// Copilot is always BYOK, even when a route does not include a provider
1032+
// credential (e.g., /_ping). Prevent the Coder token from being forwarded
1033+
// to Copilot as a provider credential.
1034+
if providerType == aibridgeconfig.ProviderCopilot {
1035+
headers.Set(agplaibridge.HeaderCoderToken, coderToken)
1036+
1037+
if extractCoderTokenFromBearerAuth(headers.Get("Authorization")) == coderToken {
1038+
headers.Del("Authorization")
1039+
}
1040+
if strings.TrimSpace(headers.Get("X-Api-Key")) == coderToken {
1041+
headers.Del("X-Api-Key")
1042+
}
1043+
return
1044+
}
1045+
1046+
// For other providers, only add the Coder token when a separate provider
1047+
// credential indicates BYOK.
1048+
if headers.Get(agplaibridge.HeaderCoderToken) != "" {
10301049
return
10311050
}
10321051

1033-
bearer := extractCoderTokenFromBearerAuth(header.Get("Authorization"))
1052+
bearer := extractCoderTokenFromBearerAuth(headers.Get("Authorization"))
10341053
if bearer != "" && bearer != coderToken {
1035-
header.Set(agplaibridge.HeaderCoderToken, coderToken)
1054+
headers.Set(agplaibridge.HeaderCoderToken, coderToken)
10361055
}
10371056
}
10381057

0 commit comments

Comments
 (0)