diff --git a/enterprise/aibridgeproxyd/aibridgeproxyd.go b/enterprise/aibridgeproxyd/aibridgeproxyd.go index dd570e7141..822161fdd6 100644 --- a/enterprise/aibridgeproxyd/aibridgeproxyd.go +++ b/enterprise/aibridgeproxyd/aibridgeproxyd.go @@ -1,12 +1,14 @@ package aibridgeproxyd import ( + "bytes" "context" "crypto/tls" "crypto/x509" "encoding/base64" "encoding/pem" "errors" + "io" "net" "net/http" "net/url" @@ -36,8 +38,14 @@ const ( // HeaderAIBridgeRequestID is the header used to correlate requests // between aibridgeproxyd and aibridged. HeaderAIBridgeRequestID = "X-AI-Bridge-Request-Id" + // ProxyAuthRealm is the realm used in Proxy-Authenticate challenges. + // The realm helps clients identify which credentials to use. + ProxyAuthRealm = `"Coder AI Bridge Proxy"` ) +// proxyAuthRequiredMsg is the response body for 407 responses. +var proxyAuthRequiredMsg = []byte(http.StatusText(http.StatusProxyAuthRequired)) + // loadMitmOnce ensures the MITM certificate is loaded exactly once. // goproxy.GoproxyCa is a package-level global variable shared across all // goproxy.ProxyHttpServer instances in the process. In tests, multiple proxy @@ -424,7 +432,9 @@ func convertDomainsToHosts(domains []string, allowedPorts []string) ([]string, e // authMiddleware is a CONNECT middleware that extracts the Coder token from // the Proxy-Authorization header and stores it in a requestContext in ctx.UserData // for use by downstream handlers. -// Requests without valid credentials are rejected. +// Requests without valid credentials receive a 407 Proxy Authentication +// Required response with a challenge header, allowing clients to retry with +// credentials. // // Clients provide credentials by setting their HTTP Proxy as: // @@ -445,12 +455,15 @@ func (s *Server) authMiddleware(host string, ctx *goproxy.ProxyCtx) (*goproxy.Co slog.F("host", host), ) - // Reject requests without valid credentials. + // Reject requests for both missing and invalid credentials if coderToken == "" { hasAuth := proxyAuth != "" logger.Warn(s.ctx, "rejecting CONNECT request", slog.F("reason", map[bool]string{true: "invalid_credentials", false: "missing_credentials"}[hasAuth]), ) + + // Send 407 challenge to allow clients to retry with credentials. + ctx.Resp = newProxyAuthRequiredResponse(ctx.Req) //nolint:bodyclose // Response body is written by goproxy to the client return goproxy.RejectConnect, host } @@ -499,6 +512,27 @@ func extractCoderTokenFromProxyAuth(proxyAuth string) string { return credentials[1] } +// newProxyAuthRequiredResponse creates a 407 Proxy Authentication Required +// response with the appropriate challenge header. This is used both during +// CONNECT handling and for decrypted requests missing authentication. +// +// Note: based on github.com/elazarl/goproxy/ext/auth.BasicUnauthorized, inlined +// here to avoid adding a dependency on the ext module. +func newProxyAuthRequiredResponse(req *http.Request) *http.Response { + return &http.Response{ + StatusCode: http.StatusProxyAuthRequired, + ProtoMajor: 1, + ProtoMinor: 1, + Request: req, + Header: http.Header{ + "Proxy-Authenticate": []string{"Basic realm=" + ProxyAuthRealm}, + "Proxy-Connection": []string{"close"}, + }, + Body: io.NopCloser(bytes.NewBuffer(proxyAuthRequiredMsg)), + ContentLength: int64(len(proxyAuthRequiredMsg)), + } +} + // defaultAIBridgeProvider maps the request host to the aibridge provider name. // - Known AI providers return their provider name, used to route to the // corresponding aibridge endpoint. @@ -563,9 +597,8 @@ func (s *Server) handleRequest(req *http.Request, ctx *goproxy.ProxyCtx) (*http. // Reject unauthenticated requests to AI providers. if reqCtx.CoderToken == "" { logger.Warn(s.ctx, "rejecting unauthenticated request to AI provider") - resp := goproxy.NewResponse(req, goproxy.ContentTypeText, http.StatusProxyAuthRequired, "Proxy authentication required") - resp.Header.Set("Proxy-Authenticate", `Basic realm="Coder AI Bridge Proxy"`) - return req, resp + // Describe to the client how to authenticate with the proxy. + return req, newProxyAuthRequiredResponse(req) } // Store provider in context for response handler. diff --git a/enterprise/aibridgeproxyd/aibridgeproxyd_test.go b/enterprise/aibridgeproxyd/aibridgeproxyd_test.go index 4db6922e6a..6081d480f3 100644 --- a/enterprise/aibridgeproxyd/aibridgeproxyd_test.go +++ b/enterprise/aibridgeproxyd/aibridgeproxyd_test.go @@ -1,6 +1,8 @@ package aibridgeproxyd_test import ( + "bufio" + "bytes" "crypto/rand" "crypto/rsa" "crypto/tls" @@ -278,6 +280,42 @@ func makeProxyAuthHeader(token string) string { return "Basic " + credentials } +// sendConnect sends a raw CONNECT request to the proxy and returns the response. +// This is needed to test proxy authentication challenges because Go's HTTP client +// doesn't expose the response when CONNECT fails with a non-2xx status. +func sendConnect(t *testing.T, proxyAddr, targetHost, proxyAuth string) *http.Response { + t.Helper() + + conn, err := net.Dial("tcp", proxyAddr) + require.NoError(t, err) + t.Cleanup(func() { _ = conn.Close() }) + + // Build CONNECT request. + var reqBuf bytes.Buffer + _, err = fmt.Fprintf(&reqBuf, "CONNECT %s HTTP/1.1\r\n", targetHost) + require.NoError(t, err) + _, err = fmt.Fprintf(&reqBuf, "Host: %s\r\n", targetHost) + require.NoError(t, err) + if proxyAuth != "" { + _, err = fmt.Fprintf(&reqBuf, "Proxy-Authorization: %s\r\n", proxyAuth) + require.NoError(t, err) + } + _, err = reqBuf.WriteString("\r\n") + require.NoError(t, err) + + // Send the CONNECT request to the proxy. + _, err = conn.Write(reqBuf.Bytes()) + require.NoError(t, err) + + // Read and parse the proxy's response. + // On success (200), the proxy establishes a tunnel. + // On auth failure (407), the proxy returns a challenge with Proxy-Authenticate header. + resp, err := http.ReadResponse(bufio.NewReader(conn), nil) + require.NoError(t, err) + + return resp +} + func TestNew(t *testing.T) { t.Parallel() @@ -777,29 +815,29 @@ func TestProxy_Authentication(t *testing.T) { t.Parallel() tests := []struct { - name string - proxyAuth string - expectError bool + name string + proxyAuth string + expectSuccess bool }{ { - name: "ValidCredentials", - proxyAuth: makeProxyAuthHeader("test-coder-token"), - expectError: false, + name: "ValidCredentials", + proxyAuth: makeProxyAuthHeader("test-coder-token"), + expectSuccess: true, }, { - name: "MissingCredentials", - proxyAuth: "", - expectError: true, + name: "MissingCredentials", + proxyAuth: "", + expectSuccess: false, }, { - name: "InvalidBase64", - proxyAuth: "Basic not-valid-base64!", - expectError: true, + name: "InvalidBase64", + proxyAuth: "Basic not-valid-base64!", + expectSuccess: false, }, { - name: "EmptyToken", - proxyAuth: makeProxyAuthHeader(""), - expectError: true, + name: "EmptyToken", + proxyAuth: makeProxyAuthHeader(""), + expectSuccess: false, }, } @@ -827,15 +865,12 @@ func TestProxy_Authentication(t *testing.T) { withDomainAllowlist(targetURL.Hostname()), ) - // Make a request through the proxy to the target server. - client := newProxyClient(t, srv, tt.proxyAuth, getProxyCertPool(t)) - req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, targetURL.String(), nil) - require.NoError(t, err) - resp, err := client.Do(req) - - if tt.expectError { - require.Error(t, err) - } else { + if tt.expectSuccess { + // Use the standard HTTP client for successful requests. + client := newProxyClient(t, srv, tt.proxyAuth, getProxyCertPool(t)) + req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, targetURL.String(), nil) + require.NoError(t, err) + resp, err := client.Do(req) require.NoError(t, err) defer resp.Body.Close() @@ -844,6 +879,25 @@ func TestProxy_Authentication(t *testing.T) { require.NoError(t, err) require.Equal(t, http.StatusOK, resp.StatusCode) require.Equal(t, "hello from aibridged", string(body)) + } else { + // Verify the proxy returns a 407 challenge with Proxy-Authenticate header. + // A raw CONNECT request is sent because Go's HTTP client doesn't expose + // the response when CONNECT fails with a non-2xx status. + resp := sendConnect(t, srv.Addr(), targetURL.Host, tt.proxyAuth) + defer resp.Body.Close() + + // Verify the status code indicates proxy authentication is required. + require.Equal(t, http.StatusProxyAuthRequired, resp.StatusCode) + + // Verify the Proxy-Authenticate header is present and contains the + // expected realm. This header tells clients how to authenticate. + proxyAuthenticate := resp.Header.Get("Proxy-Authenticate") + require.Equal(t, "Basic realm="+aibridgeproxyd.ProxyAuthRealm, proxyAuthenticate) + + // Verify the response body contains the expected error message. + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + require.Equal(t, http.StatusText(http.StatusProxyAuthRequired), string(body)) } }) }