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

Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
20 changes: 20 additions & 0 deletions aibridge/bridge.go
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@ import (
"github.com/sony/gobreaker/v2"
"go.opentelemetry.io/otel/codes"
"go.opentelemetry.io/otel/trace"
"golang.org/x/net/http/httpguts"
"golang.org/x/xerrors"

"cdr.dev/slog/v3"
Expand Down Expand Up @@ -249,6 +250,18 @@ func newInterceptionProcessor(p provider.Provider, cbs *circuitbreaker.ProviderC
client := GuessClient(r)
sessionID := GuessSessionID(client, r)

if isWebSocketUpgrade(r) {
route := strings.TrimPrefix(r.URL.Path, fmt.Sprintf("/%s", p.Name()))
logger.Debug(ctx, "rejecting unsupported WebSocket upgrade",
slog.F("provider", p.Name()),
slog.F("route", route),
slog.F("client", string(client)),
slog.F("client_session_id", sessionID),
)
http.Error(w, "WebSocket transport is not supported, use HTTP", http.StatusNotImplemented)
return
}

// Read and validate Agent Firewall correlation headers. The
// values are captured here and recorded below; the headers
// themselves are stripped from the upstream request by
Expand Down Expand Up @@ -382,6 +395,13 @@ func newInterceptionProcessor(p provider.Provider, cbs *circuitbreaker.ProviderC
}
}

// isWebSocketUpgrade reports whether r is a WebSocket opening handshake.
func isWebSocketUpgrade(r *http.Request) bool {
return r.Method == http.MethodGet &&
httpguts.HeaderValuesContainsToken(r.Header.Values("Connection"), "upgrade") &&
httpguts.HeaderValuesContainsToken(r.Header.Values("Upgrade"), "websocket")
}

// writeRequestBodyTooLarge writes a human-readable 413 response indicating that
// the request body exceeded maxRequestBodyBytes.
//
Expand Down
30 changes: 30 additions & 0 deletions aibridge/bridge_internal_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,36 @@ import (
agplaibridge "github.com/coder/coder/v2/coderd/aibridge"
)

func TestIsWebSocketUpgrade(t *testing.T) {
t.Parallel()

tests := []struct {
name string
method string
connection string
upgrade string
want bool
}{
{name: "websocket upgrade", method: http.MethodGet, connection: "keep-alive, Upgrade", upgrade: "WebSocket", want: true},
{name: "non-GET request", method: http.MethodPost, connection: "Upgrade", upgrade: "websocket", want: false},
{name: "missing connection upgrade", method: http.MethodGet, connection: "keep-alive", upgrade: "websocket", want: false},
{name: "different upgrade protocol", method: http.MethodGet, connection: "Upgrade", upgrade: "h2c", want: false},
}

for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()

req, err := http.NewRequestWithContext(t.Context(), tc.method, "/", nil)
require.NoError(t, err)
req.Header.Set("Connection", tc.connection)
req.Header.Set("Upgrade", tc.upgrade)

assert.Equal(t, tc.want, isWebSocketUpgrade(req))
})
}
}

func TestExtractAgentFirewallHeaders(t *testing.T) {
t.Parallel()

Expand Down
63 changes: 57 additions & 6 deletions aibridge/bridge_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -13,11 +13,13 @@ import (
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"go.opentelemetry.io/otel"
"go.opentelemetry.io/otel/trace"

"cdr.dev/slog/v3/sloggers/slogtest"
"github.com/coder/coder/v2/aibridge"
"github.com/coder/coder/v2/aibridge/aibridgetest"
"github.com/coder/coder/v2/aibridge/config"
"github.com/coder/coder/v2/aibridge/intercept"
"github.com/coder/coder/v2/aibridge/internal/testutil"
"github.com/coder/coder/v2/aibridge/provider"
"github.com/coder/coder/v2/coderd/httpapi"
Expand Down Expand Up @@ -187,11 +189,12 @@ func TestPassthroughRoutesForProviders(t *testing.T) {

upstreamRespBody := "upstream response"
tests := []struct {
name string
baseURLPath string
requestPath string
provider func(*testing.T, string) provider.Provider
expectPath string
name string
baseURLPath string
requestMethod string
Comment thread
ssncferreira marked this conversation as resolved.
requestPath string
provider func(*testing.T, string) provider.Provider
expectPath string
}{
{
name: "openAI_no_base_path",
Expand Down Expand Up @@ -244,6 +247,23 @@ func TestPassthroughRoutesForProviders(t *testing.T) {
},
expectPath: "/v1/models",
},
{
name: "copilot_ping",
requestPath: "/copilot/_ping",
provider: func(_ *testing.T, baseURL string) provider.Provider {
return aibridge.NewCopilotProvider(config.Copilot{BaseURL: baseURL})
},
expectPath: "/_ping",
},
{
name: "copilot_auto",
requestMethod: http.MethodPost,
requestPath: "/copilot/auto",
provider: func(_ *testing.T, baseURL string) provider.Provider {
return aibridge.NewCopilotProvider(config.Copilot{BaseURL: baseURL})
},
expectPath: "/auto",
},
}

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

req := httptest.NewRequest("", tc.requestPath, nil)
req := httptest.NewRequest(tc.requestMethod, tc.requestPath, nil)
resp := httptest.NewRecorder()
bridge.ServeHTTP(resp, req)

Expand All @@ -274,6 +294,37 @@ func TestPassthroughRoutesForProviders(t *testing.T) {
}
}

func TestWebSocketUpgradeRejected(t *testing.T) {
t.Parallel()

interceptorCalled := false
prov := &testutil.MockProvider{
NameStr: "test",
Bridged: []string{"/responses"},
InterceptorFunc: func(http.ResponseWriter, *http.Request, trace.Tracer) (intercept.Interceptor, error) {
interceptorCalled = true
return nil, nil //nolint:nilnil // The interceptor must not be reached.
},
}
bridge, err := aibridge.NewRequestBridge(
t.Context(),
[]provider.Provider{prov},
nil, nil, slogtest.Make(t, nil), nil, bridgeTestTracer,
)
require.NoError(t, err)

req := httptest.NewRequest(http.MethodGet, "/test/responses", nil)
req.Header.Set("Connection", "keep-alive, Upgrade")
req.Header.Set("Upgrade", "WebSocket")
resp := httptest.NewRecorder()

bridge.ServeHTTP(resp, req)

assert.Equal(t, http.StatusNotImplemented, resp.Code)
assert.Contains(t, resp.Body.String(), "WebSocket transport is not supported, use HTTP")
assert.False(t, interceptorCalled)
}

func TestRequestBodySizeLimit(t *testing.T) {
t.Parallel()

Expand Down
2 changes: 2 additions & 0 deletions aibridge/provider/copilot.go
Original file line number Diff line number Diff line change
Expand Up @@ -91,6 +91,8 @@ func (*Copilot) BridgedRoutes() []string {

func (*Copilot) PassthroughRoutes() []string {
return []string{
"/_ping",
"/auto",
"/models",
"/models/",
"/agents/",
Expand Down
71 changes: 45 additions & 26 deletions enterprise/aibridgeproxyd/aibridgeproxyd.go
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@ import (
"golang.org/x/xerrors"

"cdr.dev/slog/v3"
aibridgeconfig "github.com/coder/coder/v2/aibridge/config"
agplaibridge "github.com/coder/coder/v2/coderd/aibridge"
)

Expand Down Expand Up @@ -132,7 +133,7 @@ type Server struct {
// refreshProviders fetches the live provider snapshot on Reload.
// Nil disables hot-reload.
refreshProviders RefreshProvidersFunc
// providerRouter holds the live (mitmHosts, nameByHost) pair.
// providerRouter holds the live routing snapshot.
providerRouter atomic.Pointer[providerRouter]
// allowedPorts is the port allowlist for CONNECT requests. Fixed at
// construction; not reloadable.
Expand All @@ -149,19 +150,26 @@ type Server struct {
metrics *Metrics
}

type routedProvider struct {
name string
providerType string
}

// providerRouter keeps CONNECT matching and provider lookup in sync.
type providerRouter struct {
mitmHosts []string // host:port set the goproxy condition matches against.
nameByHost map[string]string // lowercase hostname -> provider name.
mitmHosts []string // host:port set the goproxy condition matches against.
providerByHost map[string]routedProvider // lowercase hostname -> provider.
}

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

func (r *providerRouter) providerFromHost(host string) string {
return r.nameByHost[strings.ToLower(host)]
func (r *providerRouter) providerFromHost(host string) routedProvider {
return r.providerByHost[strings.ToLower(host)]
}

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

logger = logger.With(
slog.F("provider", provider),
slog.F("provider", provider.name),
)

proxyAuth := ctx.Req.Header.Get("Proxy-Authorization")
Expand All @@ -685,7 +693,7 @@ func (s *Server) authMiddleware(host string, ctx *goproxy.ProxyCtx) (*goproxy.Co
ctx.UserData = &requestContext{
ConnectSessionID: connectSessionID,
CoderToken: coderToken,
Provider: provider,
Provider: provider.name,
}

logger.Debug(s.ctx, "request CONNECT authenticated")
Expand Down Expand Up @@ -936,14 +944,14 @@ func (s *Server) handleRequest(req *http.Request, ctx *goproxy.ProxyCtx) (*http.
}
}
liveProvider := s.loadProviderRouter().providerFromHost(host)
if liveProvider == "" || liveProvider != reqCtx.Provider {
if liveProvider.name == "" || liveProvider.name != reqCtx.Provider {
s.logger.Warn(s.ctx, "provider mapping changed or removed since CONNECT, passing through",
slog.F("connect_id", reqCtx.ConnectSessionID.String()),
slog.F("host", req.Host),
slog.F("method", req.Method),
slog.F("path", originalPath),
slog.F("connect_provider", reqCtx.Provider),
slog.F("live_provider", liveProvider),
slog.F("live_provider", liveProvider.name),
)
return req, nil
}
Expand Down Expand Up @@ -992,7 +1000,8 @@ func (s *Server) handleRequest(req *http.Request, ctx *goproxy.ProxyCtx) (*http.
req.URL = parsedGatewayTargetURL
req.Host = parsedGatewayTargetURL.Host

injectBYOKHeaderIfNeeded(req.Header, reqCtx.CoderToken)
// Prepare Coder authentication for centralized and BYOK requests.
prepareAIGatewayAuth(req.Header, reqCtx.CoderToken, liveProvider.providerType)

// Set request ID header to correlate requests between aibridgeproxyd and aibridged.
req.Header.Set(agplaibridge.HeaderCoderRequestID, reqCtx.RequestID.String())
Expand All @@ -1019,24 +1028,34 @@ func (s *Server) handleRequest(req *http.Request, ctx *goproxy.ProxyCtx) (*http.
return req, nil
}

// injectBYOKHeaderIfNeeded sets HeaderCoderToken when the
// Authorization header carries a bearer token that differs from the
// Coder token, indicating the client is using its own LLM
// credentials. Clients that can set custom headers
// do this themselves; this handles clients that cannot.
//
// In centralized mode, Authorization carries the Coder token
// itself, so aibridged discovers it via ExtractAuthToken
// without any extra header.
func injectBYOKHeaderIfNeeded(header http.Header, coderToken string) {
// Don’t overwrite the header if it’s already set.
if header.Get(agplaibridge.HeaderCoderToken) != "" {
// prepareAIGatewayAuth prepares the Coder authentication headers for AI
// Gateway. Copilot is always BYOK, while other providers may use centralized
// or BYOK authentication.
func prepareAIGatewayAuth(headers http.Header, coderToken, providerType string) {
// Copilot is always BYOK, even when a route does not include a provider
// credential (e.g., /_ping). Prevent the Coder token from being forwarded
// to Copilot as a provider credential.
if providerType == aibridgeconfig.ProviderCopilot {
headers.Set(agplaibridge.HeaderCoderToken, coderToken)

if extractCoderTokenFromBearerAuth(headers.Get("Authorization")) == coderToken {
headers.Del("Authorization")
}
if strings.TrimSpace(headers.Get("X-Api-Key")) == coderToken {
headers.Del("X-Api-Key")
}
return
}

// For other providers, only add the Coder token when a separate provider
// credential indicates BYOK.
if headers.Get(agplaibridge.HeaderCoderToken) != "" {
return
}

bearer := extractCoderTokenFromBearerAuth(header.Get("Authorization"))
bearer := extractCoderTokenFromBearerAuth(headers.Get("Authorization"))
if bearer != "" && bearer != coderToken {
header.Set(agplaibridge.HeaderCoderToken, coderToken)
headers.Set(agplaibridge.HeaderCoderToken, coderToken)
}
}

Expand Down
Loading
Loading