diff --git a/enterprise/aibridgeproxyd/aibridgeproxyd.go b/enterprise/aibridgeproxyd/aibridgeproxyd.go index 8886d9b0a3..dd570e7141 100644 --- a/enterprise/aibridgeproxyd/aibridgeproxyd.go +++ b/enterprise/aibridgeproxyd/aibridgeproxyd.go @@ -18,6 +18,7 @@ import ( "github.com/elazarl/goproxy" "github.com/go-chi/chi/v5" + "github.com/google/uuid" "golang.org/x/xerrors" "cdr.dev/slog/v3" @@ -31,6 +32,12 @@ const ( HostOpenAI = "api.openai.com" ) +const ( + // HeaderAIBridgeRequestID is the header used to correlate requests + // between aibridgeproxyd and aibridged. + HeaderAIBridgeRequestID = "X-AI-Bridge-Request-Id" +) + // 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 @@ -56,6 +63,26 @@ type Server struct { caCert []byte } +// requestContext holds metadata propagated through the proxy request/response chain. +// It is stored in goproxy's ProxyCtx.UserData and enriched as the request progresses +// through the proxy handlers. +type requestContext struct { + // ConnectSessionID is a unique identifier for this CONNECT session. + // Set in authMiddleware during the CONNECT handshake. + // Used to correlate requests/responses with their originating CONNECT. + ConnectSessionID uuid.UUID + // CoderToken is the authentication token extracted from Proxy-Authorization. + // Set in authMiddleware during the CONNECT handshake. + CoderToken string + // RequestID is a unique identifier for this request. + // Set in handleRequest for MITM'd requests. + // Sent to aibridged via custom header for cross-service correlation. + RequestID uuid.UUID + // Provider is the aibridge provider name. + // Set in handleRequest when handling MITM requests for allowlisted domains. + Provider string +} + // Options configures the AI Bridge Proxy server. type Options struct { // ListenAddr is the address the proxy server will listen on. @@ -93,7 +120,7 @@ type Options struct { } func New(ctx context.Context, logger slog.Logger, opts Options) (*Server, error) { - logger.Info(ctx, "initializing AI Bridge Proxy server") + logger.Info(ctx, "initializing aibridgeproxyd") if opts.ListenAddr == "" { return nil, xerrors.New("listen address is required") @@ -140,11 +167,6 @@ func New(ctx context.Context, logger slog.Logger, opts Options) (*Server, error) } } - logger.Info(ctx, "configured domain allowlist for MITM", - slog.F("domains", opts.DomainAllowlist), - slog.F("hosts", mitmHosts), - ) - // Load CA certificate for MITM certPEM, err := loadMitmCertificate(opts.CertFile, opts.KeyFile) if err != nil { @@ -182,10 +204,6 @@ func New(ctx context.Context, logger slog.Logger, opts Options) (*Server, error) return nil, xerrors.Errorf("invalid upstream proxy URL %q: %w", opts.UpstreamProxy, err) } - logger.Info(ctx, "configuring upstream proxy for tunneled requests", - slog.F("upstream", upstreamURL.Host), - ) - // Set transport without Proxy to ensure MITM'd requests go directly to aibridge, // not through any upstream proxy. proxy.Tr = &http.Transport{ @@ -244,6 +262,8 @@ func New(ctx context.Context, logger slog.Logger, opts Options) (*Server, error) // Handle decrypted requests: route to aibridged for known AI providers, or tunnel to original destination. proxy.OnRequest().DoFunc(srv.handleRequest) + // Handle responses from aibridged. + proxy.OnResponse().DoFunc(srv.handleResponse) // Create listener first so we can get the actual address. // This is useful in tests where port 0 is used to avoid conflicts. @@ -259,8 +279,15 @@ func New(ctx context.Context, logger slog.Logger, opts Options) (*Server, error) ReadHeaderTimeout: 10 * time.Second, } + logger.Info(ctx, "aibridgeproxyd configured", + slog.F("listen_addr", listener.Addr().String()), + slog.F("coder_access_url", coderAccessURL.String()), + slog.F("domain_allowlist", mitmHosts), + slog.F("upstream_proxy", opts.UpstreamProxy), + ) + go func() { - logger.Info(ctx, "starting AI Bridge Proxy", slog.F("addr", listener.Addr().String())) + logger.Info(ctx, "starting aibridgeproxyd server", slog.F("addr", listener.Addr().String())) if err := srv.httpServer.Serve(listener); err != nil && !errors.Is(err, http.ErrServerClosed) { logger.Error(ctx, "aibridgeproxyd server error", slog.Error(err)) } @@ -283,6 +310,7 @@ func (s *Server) Close() error { if s.httpServer == nil { return nil } + s.logger.Info(s.ctx, "closing aibridgeproxyd server") ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() return s.httpServer.Shutdown(ctx) @@ -334,26 +362,26 @@ func (s *Server) portMiddleware(allowedPorts []string) func(host string, ctx *go } return func(host string, _ *goproxy.ProxyCtx) (*goproxy.ConnectAction, string) { + logger := s.logger.With( + slog.F("host", host), + ) + _, port, err := net.SplitHostPort(host) if err != nil { - s.logger.Warn(s.ctx, "rejecting CONNECT with invalid host format", - slog.F("host", host), + logger.Warn(s.ctx, "rejecting CONNECT with invalid host format", slog.Error(err), ) return goproxy.RejectConnect, host } if port == "" { - s.logger.Warn(s.ctx, "rejecting CONNECT with empty port", - slog.F("host", host), - ) + logger.Warn(s.ctx, "rejecting CONNECT with empty port") return goproxy.RejectConnect, host } + logger = logger.With(slog.F("port", port)) + if !allowed[port] { - s.logger.Warn(s.ctx, "rejecting CONNECT to non-allowed port", - slog.F("host", host), - slog.F("port", port), - ) + logger.Warn(s.ctx, "rejecting CONNECT to non-allowed port") return goproxy.RejectConnect, host } @@ -394,8 +422,8 @@ 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 ctx.UserData for use by -// downstream request handlers. +// the Proxy-Authorization header and stores it in a requestContext in ctx.UserData +// for use by downstream handlers. // Requests without valid credentials are rejected. // // Clients provide credentials by setting their HTTP Proxy as: @@ -404,23 +432,37 @@ func convertDomainsToHosts(domains []string, allowedPorts []string) ([]string, e // // The token is extracted from the password field of basic auth. func (s *Server) authMiddleware(host string, ctx *goproxy.ProxyCtx) (*goproxy.ConnectAction, string) { + // Generate a unique connect session ID for this CONNECT request. + // A UUID is used instead of goproxy's ctx.Session because ctx.Session is an + // incrementing int64 that resets on process restart and is not globally unique. + connectSessionID := uuid.New() + proxyAuth := ctx.Req.Header.Get("Proxy-Authorization") coderToken := extractCoderTokenFromProxyAuth(proxyAuth) + logger := s.logger.With( + slog.F("connect_id", connectSessionID), + slog.F("host", host), + ) + // Reject requests without valid credentials. if coderToken == "" { hasAuth := proxyAuth != "" - s.logger.Warn(s.ctx, "rejecting CONNECT request", - slog.F("host", host), + logger.Warn(s.ctx, "rejecting CONNECT request", slog.F("reason", map[bool]string{true: "invalid_credentials", false: "missing_credentials"}[hasAuth]), ) return goproxy.RejectConnect, host } - // Store the token in UserData for downstream handlers. - // goproxy propagates UserData to subsequent request contexts + // Store the request context in UserData for downstream handlers. + // goproxy propagates UserData to subsequent request/response contexts // for decrypted requests within this MITM session. - ctx.UserData = coderToken + ctx.UserData = &requestContext{ + ConnectSessionID: connectSessionID, + CoderToken: coderToken, + } + + logger.Debug(s.ctx, "request CONNECT authenticated") return goproxy.MitmConnect, host } @@ -479,6 +521,31 @@ func defaultAIBridgeProvider(host string) string { func (s *Server) handleRequest(req *http.Request, ctx *goproxy.ProxyCtx) (*http.Request, *http.Response) { originalPath := req.URL.Path + // Get the request context stored during CONNECT. + reqCtx, _ := ctx.UserData.(*requestContext) + if reqCtx == nil { + s.logger.Warn(s.ctx, "rejecting request with missing context", + slog.F("host", req.Host), + slog.F("method", req.Method), + slog.F("path", originalPath), + ) + 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 + } + + // Generate a unique request ID for this request. + // This ID is sent to aibridged for cross-service log correlation. + reqCtx.RequestID = uuid.New() + + logger := s.logger.With( + slog.F("connect_id", reqCtx.ConnectSessionID.String()), + slog.F("request_id", reqCtx.RequestID.String()), + slog.F("host", req.Host), + slog.F("method", req.Method), + slog.F("path", originalPath), + ) + // Check if this request is for a supported AI provider. provider := s.aibridgeProviderFromHost(req.URL.Hostname()) if provider == "" { @@ -487,44 +554,38 @@ func (s *Server) handleRequest(req *http.Request, ctx *goproxy.ProxyCtx) (*http. // The request is MITM'd (decrypted) but since there is no mapping, // there is no known route to aibridge. // Log error and forward to the original destination as a fallback. - s.logger.Error(s.ctx, "decrypted request has no provider mapping, passing through", - slog.F("host", req.Host), - slog.F("method", req.Method), - slog.F("path", originalPath), - ) + logger.Error(s.ctx, "decrypted request has no provider mapping, passing through") return req, nil } - // Get the Coder token stored during CONNECT. - coderToken, _ := ctx.UserData.(string) + logger = logger.With(slog.F("provider", provider)) // Reject unauthenticated requests to AI providers. - if coderToken == "" { - s.logger.Warn(s.ctx, "rejecting unauthenticated request to AI provider", - slog.F("host", req.Host), - slog.F("provider", provider), - ) + if reqCtx.CoderToken == "" { + logger.Warn(s.ctx, "rejecting unauthenticated request to AI provider") resp := goproxy.NewResponse(req, goproxy.ContentTypeText, http.StatusProxyAuthRequired, "Proxy authentication required") - // Describe to the client how to authenticate with the proxy. resp.Header.Set("Proxy-Authenticate", `Basic realm="Coder AI Bridge Proxy"`) return req, resp } + // Store provider in context for response handler. + reqCtx.Provider = provider + // Rewrite the request to point to aibridged. if s.coderAccessURL == nil || s.coderAccessURL.String() == "" { - s.logger.Error(s.ctx, "coderAccessURL is not configured") + logger.Error(s.ctx, "coderAccessURL is not configured") return req, goproxy.NewResponse(req, goproxy.ContentTypeText, http.StatusInternalServerError, "Proxy misconfigured") } aiBridgeURL, err := url.JoinPath(s.coderAccessURL.String(), "api/v2/aibridge", provider, originalPath) if err != nil { - s.logger.Error(s.ctx, "failed to build aibridged URL", slog.Error(err)) + logger.Error(s.ctx, "failed to build aibridged URL", slog.Error(err)) return req, goproxy.NewResponse(req, goproxy.ContentTypeText, http.StatusInternalServerError, "Failed to build AI Bridge URL") } aiBridgeParsedURL, err := url.Parse(aiBridgeURL) if err != nil { - s.logger.Error(s.ctx, "failed to parse aibridged URL", slog.Error(err)) + logger.Error(s.ctx, "failed to parse aibridged URL", slog.Error(err)) return req, goproxy.NewResponse(req, goproxy.ContentTypeText, http.StatusInternalServerError, "Failed to parse AI Bridge URL") } @@ -537,17 +598,47 @@ func (s *Server) handleRequest(req *http.Request, ctx *goproxy.ProxyCtx) (*http. // Set X-Coder-Token header for aibridged authentication. // Using a separate header preserves the original request headers, // which are forwarded to upstream providers. - req.Header.Set(agplaibridge.HeaderCoderAuth, coderToken) + req.Header.Set(agplaibridge.HeaderCoderAuth, reqCtx.CoderToken) - s.logger.Debug(s.ctx, "routing request to aibridged", - slog.F("provider", provider), - slog.F("original_path", originalPath), + // Set custom header for cross-service log correlation. + // This allows correlating aibridgeproxyd logs with aibridged logs. + req.Header.Set(HeaderAIBridgeRequestID, reqCtx.RequestID.String()) + + logger.Info(s.ctx, "routing MITM request to aibridged", slog.F("aibridged_url", aiBridgeParsedURL.String()), ) return req, nil } +// handleResponse handles responses received from aibridged. +// This is only called for MITM'd requests (allowlisted domains routed through aibridged). +// Tunneled requests (non-allowlisted domains) bypass this handler entirely. +func (s *Server) handleResponse(resp *http.Response, ctx *goproxy.ProxyCtx) *http.Response { + if resp == nil { + return nil + } + + reqCtx, _ := ctx.UserData.(*requestContext) + connectSessionID := uuid.Nil + requestID := uuid.Nil + provider := "" + if reqCtx != nil { + connectSessionID = reqCtx.ConnectSessionID + requestID = reqCtx.RequestID + provider = reqCtx.Provider + } + + s.logger.Debug(s.ctx, "received response from aibridged", + slog.F("connect_id", connectSessionID.String()), + slog.F("request_id", requestID.String()), + slog.F("status", resp.StatusCode), + slog.F("provider", provider), + ) + + return resp +} + // Handler returns an HTTP handler for the AI Bridge Proxy's HTTP endpoints. // This is separate from the proxy server itself and is used by coderd to // serve endpoints like the CA certificate. diff --git a/enterprise/aibridgeproxyd/aibridgeproxyd_test.go b/enterprise/aibridgeproxyd/aibridgeproxyd_test.go index 3c0103a303..4db6922e6a 100644 --- a/enterprise/aibridgeproxyd/aibridgeproxyd_test.go +++ b/enterprise/aibridgeproxyd/aibridgeproxyd_test.go @@ -22,9 +22,11 @@ import ( "testing" "time" + "github.com/google/uuid" "github.com/stretchr/testify/require" "golang.org/x/xerrors" + "cdr.dev/slog/v3" "cdr.dev/slog/v3/sloggers/slogtest" agplaibridge "github.com/coder/coder/v2/coderd/aibridge" "github.com/coder/coder/v2/enterprise/aibridgeproxyd" @@ -171,7 +173,7 @@ func newTestProxy(t *testing.T, opts ...testProxyOption) *aibridgeproxyd.Server } certFile, keyFile := getSharedTestCA(t) - logger := slogtest.Make(t, nil) + logger := slogtest.Make(t, nil).Leveled(slog.LevelDebug) aibridgeOpts := aibridgeproxyd.Options{ ListenAddr: cfg.listenAddr, @@ -910,12 +912,13 @@ func TestProxy_MITM(t *testing.T) { t.Parallel() // Track what aibridged receives. - var receivedPath, receivedCoderToken string + var receivedPath, receivedCoderToken, receivedRequestID string // Create a mock aibridged server that captures requests. aibridgedServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { receivedPath = r.URL.Path receivedCoderToken = r.Header.Get(agplaibridge.HeaderCoderAuth) + receivedRequestID = r.Header.Get(aibridgeproxyd.HeaderAIBridgeRequestID) w.WriteHeader(http.StatusOK) _, _ = w.Write([]byte("hello from aibridged")) })) @@ -984,11 +987,15 @@ func TestProxy_MITM(t *testing.T) { require.Equal(t, "hello from tunneled", string(body)) require.Empty(t, receivedPath, "aibridged should not receive tunneled requests") require.Empty(t, receivedCoderToken, "tunneled requests are not authenticated by the proxy") + require.Empty(t, receivedRequestID, "tunneled requests should not have request ID header") } else { // Verify the request was routed to aibridged correctly. require.Equal(t, "hello from aibridged", string(body)) require.Equal(t, tt.expectedPath, receivedPath) require.Equal(t, "test-token", receivedCoderToken, "MITM'd requests must include Coder token") + require.NotEmpty(t, receivedRequestID, "MITM'd requests must include request ID header") + _, err := uuid.Parse(receivedRequestID) + require.NoError(t, err, "request ID must be a valid UUID") } }) }