diff --git a/cli/root.go b/cli/root.go index ed89a00ddce38..fc20141dc1583 100644 --- a/cli/root.go +++ b/cli/root.go @@ -58,6 +58,8 @@ var ( // anything. ErrSilent = xerrors.New("silent error") + ErrClientURLNotConfigured = xerrors.New("client URL is not configured") + errKeyringNotSupported = xerrors.New("keyring storage is not supported on this operating system; omit --use-keyring to use file-based storage") ) @@ -602,23 +604,58 @@ func (r *RootCmd) SetClock(clk quartz.Clock) { // ensureClientURL loads the client URL from the config file if it // wasn't provided via --url or CODER_URL. func (r *RootCmd) ensureClientURL() error { - if r.clientURL != nil && r.clientURL.String() != "" { - return nil - } - rawURL, err := r.createConfig().URL().Read() - // If the configuration files are absent, the user is logged out. - if os.IsNotExist(err) { - binPath, err := os.Executable() - if err != nil { + u, err := r.resolveClientURL() + + if errors.Is(err, ErrClientURLNotConfigured) { + binPath, execErr := os.Executable() + if execErr != nil { binPath = "coder" } return xerrors.Errorf(notLoggedInMessage, binPath) } + if err != nil { return err } - r.clientURL, err = url.Parse(strings.TrimSpace(rawURL)) - return err + + r.clientURL = u + return nil +} + +func (r *RootCmd) resolveClientURL() (*url.URL, error) { + if r.clientURL != nil && r.clientURL.String() != "" { + return r.clientURL, nil + } + + rawURL, err := r.createConfig().URL().Read() + if err != nil { + if errors.Is(err, os.ErrNotExist) { + return nil, ErrClientURLNotConfigured + } + return nil, xerrors.Errorf("read configured URL: %w", err) + } + parsedURL, err := url.Parse(strings.TrimSpace(rawURL)) + if err != nil { + return nil, xerrors.Errorf("parse configured URL: %w", err) + } + return parsedURL, nil +} + +// ResolveClientConnection resolves the deployment URL and client TLS transport +// without reading or requiring a user session. +func (r *RootCmd) ResolveClientConnection() (*url.URL, http.RoundTripper, error) { + serverURL, err := r.resolveClientURL() + if err != nil { + return nil, nil, err + } + if err := r.ensureTLSConfig(); err != nil { + return nil, nil, xerrors.Errorf("load client TLS config: %w", err) + } + transport, err := newHTTPTransport(r.tlsConfig) + if err != nil { + return nil, nil, xerrors.Errorf("create HTTP transport: %w", err) + } + return serverURL, transport, nil } // ensureTLSConfig loads the TLS configuration from files if specified. diff --git a/cli/root_test.go b/cli/root_test.go index fa65cf2973025..534bf9b9cf244 100644 --- a/cli/root_test.go +++ b/cli/root_test.go @@ -18,6 +18,7 @@ import ( "github.com/coder/coder/v2/buildinfo" "github.com/coder/coder/v2/cli" "github.com/coder/coder/v2/cli/clitest" + "github.com/coder/coder/v2/cli/config" "github.com/coder/coder/v2/coderd" "github.com/coder/coder/v2/coderd/coderdtest" "github.com/coder/coder/v2/codersdk" @@ -105,6 +106,148 @@ func TestCommandHelp(t *testing.T) { )) } +func TestResolveClientConnection(t *testing.T) { + t.Parallel() + + run := func(t *testing.T, configure func(config.Root), args ...string) (string, http.RoundTripper, error, error) { + t.Helper() + + var root cli.RootCmd + var gotURL string + var gotTransport http.RoundTripper + var gotErr error + cmd, err := root.Command([]*serpent.Command{{ + Use: "resolve", + Handler: func(*serpent.Invocation) error { + serverURL, transport, err := root.ResolveClientConnection() + if serverURL != nil { + gotURL = serverURL.String() + } + gotTransport = transport + gotErr = err + return nil + }, + }}) + require.NoError(t, err) + + inv, cfg := clitest.NewWithCommand(t, cmd, args...) + if configure != nil { + configure(cfg) + } + runErr := inv.Run() + return gotURL, gotTransport, gotErr, runErr + } + + tests := []struct { + name string + args []string + configure func(*testing.T, config.Root) + wantURL string + wantTransport bool + wantErr string + wantRunErr string + checkTransport func(*testing.T, http.RoundTripper) + }{ + { + name: "MissingURL", + args: []string{"resolve"}, + wantErr: cli.ErrClientURLNotConfigured.Error(), + }, + { + name: "URLFlag", + args: []string{"--url", "https://example.com", "resolve"}, + wantURL: "https://example.com", + wantTransport: true, + }, + { + name: "ConfiguredURL", + args: []string{"resolve"}, + configure: func(t *testing.T, cfg config.Root) { + t.Helper() + require.NoError(t, cfg.URL().Write("https://configured.example.com")) + }, + wantURL: "https://configured.example.com", + wantTransport: true, + }, + { + name: "URLFlagOverridesConfig", + args: []string{"--url", "https://flag.example.com", "resolve"}, + configure: func(t *testing.T, cfg config.Root) { + t.Helper() + require.NoError(t, cfg.URL().Write("https://configured.example.com")) + }, + wantURL: "https://flag.example.com", + wantTransport: true, + }, + { + name: "InvalidURLFlag", + args: []string{"--url", "%zz", "resolve"}, + wantRunErr: "invalid URL escape", + }, + { + name: "ClientTLSConfig", + args: func() []string { + certPath, keyPath := generateTLSCertificate(t) + return []string{ + "--url", "https://example.com", + "--client-tls-cert-file", certPath, + "--client-tls-key-file", keyPath, + "resolve", + } + }(), + wantURL: "https://example.com", + wantTransport: true, + checkTransport: func(t *testing.T, transport http.RoundTripper) { + t.Helper() + + httpTransport, ok := transport.(*http.Transport) + require.True(t, ok) + require.NotNil(t, httpTransport.TLSClientConfig) + require.Len(t, httpTransport.TLSClientConfig.Certificates, 1) + }, + }, + { + name: "TLSConfigError", + args: []string{ + "--url", "https://example.com", + "--client-tls-cert-file", "/tmp/missing-cert.pem", + "resolve", + }, + wantErr: "load client TLS config: --client-tls-cert-file and --client-tls-key-file must be specified together", + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + var configure func(config.Root) + if tc.configure != nil { + configure = func(cfg config.Root) { + tc.configure(t, cfg) + } + } + + serverURL, transport, err, runErr := run(t, configure, tc.args...) + if tc.wantRunErr != "" { + require.ErrorContains(t, runErr, tc.wantRunErr) + return + } + require.NoError(t, runErr) + if tc.wantErr != "" { + require.ErrorContains(t, err, tc.wantErr) + } else { + require.NoError(t, err) + } + require.Equal(t, tc.wantURL, serverURL) + require.Equal(t, tc.wantTransport, transport != nil) + if tc.checkTransport != nil { + tc.checkTransport(t, transport) + } + }) + } +} + func TestRoot(t *testing.T) { t.Parallel() t.Run("MissingRootCommand", func(t *testing.T) { diff --git a/coderd/aibridged/aibridged.go b/coderd/aibridged/aibridged.go index d70a50784722a..6a300c350c023 100644 --- a/coderd/aibridged/aibridged.go +++ b/coderd/aibridged/aibridged.go @@ -16,7 +16,11 @@ import ( "github.com/coder/retry" ) -var _ io.Closer = &Server{} +var ( + _ io.Closer = &Server{} + + ErrShutdown = xerrors.New("aibridged server shutdown") +) // Server provides the AI Bridge functionality. // It is responsible for: @@ -42,10 +46,11 @@ type Server struct { initConnectionCh chan struct{} initConnectionOnce sync.Once - // lifecycleCtx is canceled when we start closing. + // lifecycleCtx is canceled when we start closing or when the + // connection loop exits permanently. lifecycleCtx context.Context - // cancelFn closes the lifecycleCtx. - cancelFn func() + // cancelFn closes the lifecycleCtx with the reason it closed. + cancelFn context.CancelCauseFunc shutdownOnce sync.Once } @@ -55,7 +60,7 @@ func New(ctx context.Context, pool Pooler, rpcDialer Dialer, logger slog.Logger, return nil, xerrors.Errorf("nil rpcDialer given") } - ctx, cancel := context.WithCancel(ctx) + ctx, cancel := context.WithCancelCause(ctx) daemon := &Server{ logger: logger, tracer: tracer, @@ -78,6 +83,11 @@ func New(ctx context.Context, pool Pooler, rpcDialer Dialer, logger slog.Logger, func (s *Server) connect() { defer s.logger.Debug(s.lifecycleCtx, "connect loop exited") defer s.wg.Done() + defer func() { + if s.lifecycleCtx.Err() == nil { + s.cancelFn(xerrors.New("connect loop exited")) + } + }() logConnect := s.logger.With(slog.F("context", "aibridged.server")).Debug // An exponential back-off occurs when the connection is failing to dial. @@ -93,13 +103,25 @@ connectLoop: client, err := s.clientDialer(s.lifecycleCtx) if err != nil { if errors.Is(err, context.Canceled) { + if s.lifecycleCtx.Err() == nil { + s.cancelFn(err) + } return } var sdkErr *codersdk.Error - // If something is wrong with our auth, stop trying to connect. - if errors.As(err, &sdkErr) && sdkErr.StatusCode() == http.StatusForbidden { - s.logger.Error(s.lifecycleCtx, "not authorized to dial coderd", slog.Error(err)) - return + // If something is wrong with configuration, stop trying to connect. + if errors.As(err, &sdkErr) { + switch sdkErr.StatusCode() { + // These statuses are terminal failures from the /api/v2/ai-gateway/serve + // handshake: wrong gateway key, incompatible API version, or entitlement failure. + case http.StatusBadRequest, http.StatusUnauthorized, http.StatusForbidden: + err = xerrors.Errorf("dial coderd: %w", err) + s.logger.Error(s.lifecycleCtx, "fatal error dialing coderd", slog.Error(err)) + s.cancelFn(err) + return + default: + err = xerrors.Errorf("unexpected HTTP response dialing coderd: %w", err) + } } if s.isShutdown() { return @@ -133,9 +155,32 @@ connectLoop: } } +// Done returns a channel that is closed when the server lifecycle ends. +// It closes on explicit shutdown and on fatal connection-loop exit. +func (s *Server) Done() <-chan struct{} { + return s.lifecycleCtx.Done() +} + +// Err returns the reason the server lifecycle ended. +func (s *Server) Err() error { + if cause := context.Cause(s.lifecycleCtx); cause != nil { + return cause + } + return s.lifecycleCtx.Err() +} + func (s *Server) Client() (DRPCClient, error) { + return s.ClientContext(context.Background()) +} + +func (s *Server) ClientContext(ctx context.Context) (DRPCClient, error) { select { - case <-s.lifecycleCtx.Done(): + case <-ctx.Done(): + return nil, ctx.Err() + case <-s.Done(): + if err := s.Err(); err != nil { + return nil, err + } return nil, xerrors.New("context closed") case client := <-s.clientCh: return client, nil @@ -170,7 +215,7 @@ func (s *Server) isShutdown() bool { func (s *Server) Shutdown(ctx context.Context) error { var err error s.shutdownOnce.Do(func() { - s.cancelFn() + s.cancelFn(ErrShutdown) // Wait for any outstanding connections to terminate. s.wg.Wait() diff --git a/coderd/aibridged/aibridged_test.go b/coderd/aibridged/aibridged_test.go index 4ef3d603e6089..5f3deca4786a4 100644 --- a/coderd/aibridged/aibridged_test.go +++ b/coderd/aibridged/aibridged_test.go @@ -6,6 +6,7 @@ import ( "io" "net/http" "net/http/httptest" + "sync/atomic" "testing" "github.com/google/uuid" @@ -40,8 +41,13 @@ func singleKeyPool(t *testing.T, name, key string) *keypool.Pool { func newTestServer(t *testing.T) (*aibridged.Server, *mock.MockDRPCClient, *mock.MockPooler) { t.Helper() + return newTestServerWithDialer(t, nil, nil) +} + +func newTestServerWithDialer(t *testing.T, dialer aibridged.Dialer, loggerOptions *slogtest.Options) (*aibridged.Server, *mock.MockDRPCClient, *mock.MockPooler) { + t.Helper() - logger := slogtest.Make(t, nil) + logger := slogtest.Make(t, loggerOptions) ctrl := gomock.NewController(t) client := mock.NewMockDRPCClient(ctrl) pool := mock.NewMockPooler(ctrl) @@ -50,12 +56,12 @@ func newTestServer(t *testing.T) (*aibridged.Server, *mock.MockDRPCClient, *mock client.EXPECT().DRPCConn().AnyTimes().Return(conn) pool.EXPECT().Shutdown(gomock.Any()).MinTimes(1).Return(nil) - srv, err := aibridged.New( - t.Context(), - pool, - func(ctx context.Context) (aibridged.DRPCClient, error) { + if dialer == nil { + dialer = func(ctx context.Context) (aibridged.DRPCClient, error) { return client, nil - }, logger, testTracer) + } + } + srv, err := aibridged.New(t.Context(), pool, dialer, logger, testTracer) require.NoError(t, err, "create new aibridged") t.Cleanup(func() { srv.Shutdown(context.Background()) @@ -79,6 +85,39 @@ func (*mockDRPCConn) NewStream(ctx context.Context, rpc string, enc drpc.Encodin return nil, nil } +func sdkError(status int, message string) error { + return codersdk.ReadBodyAsError(&http.Response{ + StatusCode: status, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(bytes.NewBufferString(`{"message":"` + message + `"}`)), + }) +} + +func TestClient_TransientDialErrorRetries(t *testing.T) { + t.Parallel() + + var calls atomic.Int32 + ctrl := gomock.NewController(t) + client := mock.NewMockDRPCClient(ctrl) + client.EXPECT().DRPCConn().AnyTimes().Return(&mockDRPCConn{}) + pool := mock.NewMockPooler(ctrl) + pool.EXPECT().Shutdown(gomock.Any()).MinTimes(1).Return(nil) + dialFc := func(context.Context) (aibridged.DRPCClient, error) { + if calls.Add(1) == 1 { + return nil, sdkError(http.StatusInternalServerError, "internal error") + } + return client, nil + } + + srv, err := aibridged.New(t.Context(), pool, dialFc, slogtest.Make(t, nil), testTracer) + require.NoError(t, err) + t.Cleanup(func() { _ = srv.Shutdown(context.Background()) }) + + _, err = srv.ClientContext(testutil.Context(t, testutil.WaitShort)) + require.NoError(t, err) + require.Equal(t, int32(2), calls.Load()) +} + func TestServeHTTP_FailureModes(t *testing.T) { t.Parallel() @@ -91,6 +130,7 @@ func TestServeHTTP_FailureModes(t *testing.T) { applyMocksFn func(client *mock.MockDRPCClient, pool *mock.MockPooler) dialerFn aibridged.Dialer contextFn func() context.Context + ignoreLogs bool expectedErr error expectedStatus int }{ @@ -127,7 +167,34 @@ func TestServeHTTP_FailureModes(t *testing.T) { expectedStatus: http.StatusForbidden, }, - // TODO: coderd connection-related failures. + // Coderd connection-related failures. + { + name: "fatal bad request dial error", + dialerFn: func(context.Context) (aibridged.DRPCClient, error) { + return nil, sdkError(http.StatusBadRequest, "bad request") + }, + ignoreLogs: true, + expectedErr: aibridged.ErrConnect, + expectedStatus: http.StatusServiceUnavailable, + }, + { + name: "fatal unauthorized dial error", + dialerFn: func(context.Context) (aibridged.DRPCClient, error) { + return nil, sdkError(http.StatusUnauthorized, "unauthorized") + }, + ignoreLogs: true, + expectedErr: aibridged.ErrConnect, + expectedStatus: http.StatusServiceUnavailable, + }, + { + name: "fatal forbidden dial error", + dialerFn: func(context.Context) (aibridged.DRPCClient, error) { + return nil, sdkError(http.StatusForbidden, "forbidden") + }, + ignoreLogs: true, + expectedErr: aibridged.ErrConnect, + expectedStatus: http.StatusServiceUnavailable, + }, // Budget-related failures. { @@ -173,7 +240,11 @@ func TestServeHTTP_FailureModes(t *testing.T) { t.Run(tc.name, func(t *testing.T) { t.Parallel() - srv, client, pool := newTestServer(t) + var loggerOptions *slogtest.Options + if tc.ignoreLogs { + loggerOptions = &slogtest.Options{IgnoreErrors: true} + } + srv, client, pool := newTestServerWithDialer(t, tc.dialerFn, loggerOptions) conn := &mockDRPCConn{} client.EXPECT().DRPCConn().AnyTimes().Return(conn) diff --git a/coderd/aibridged/dialer.go b/coderd/aibridged/dialer.go new file mode 100644 index 0000000000000..6b2a17d7c17da --- /dev/null +++ b/coderd/aibridged/dialer.go @@ -0,0 +1,101 @@ +package aibridged + +import ( + "context" + "errors" + "io" + "net/http" + "net/url" + + "github.com/hashicorp/yamux" + "golang.org/x/xerrors" + + "github.com/coder/coder/v2/buildinfo" + aibridgedproto "github.com/coder/coder/v2/coderd/aibridged/proto" + "github.com/coder/coder/v2/codersdk" + "github.com/coder/coder/v2/codersdk/drpcsdk" + "github.com/coder/websocket" +) + +// NewWebsocketDialer returns a [Dialer] that connects a standalone AI +// Gateway to coderd's /api/v2/ai-gateway/serve endpoint over a WebSocket, +// multiplexes it with yamux, and exposes the aibridged DRPC services +// (Recorder, MCPConfigurator, Authorizer, ProviderConfigurator) over it. +// This is the standalone counterpart to API.CreateInMemoryAIBridgeServer, +// which wires the same services over an in-memory pipe for the embedded +// daemon. +// +// The gateway authenticates with an AI Gateway key +// (codersdk.AIGatewayKeyHeader), advertises its API version via the +// "version" query parameter, and reports its build version via +// codersdk.BuildVersionHeader (used by coderd for observability only). +// TLS for this connection is governed by the scheme of serverURL and any +// TLS configuration baked into transport. +// +// On a failed upgrade the coderd HTTP error is returned as a +// *codersdk.Error so [Server.connect] can distinguish fatal +// auth/entitlement failures from transient ones. +func readAIGatewayServeError(res *http.Response) error { + err := codersdk.ReadBodyAsError(res) + + var sdkErr *codersdk.Error + if errors.As(err, &sdkErr) && res.StatusCode == http.StatusUnauthorized { + // /ai-gateway/serve authenticates with an AI Gateway key, not a user + // session. Generic user-login helpers are misleading here. + sdkErr.Helper = "" + } + return err +} + +func NewWebsocketDialer(serverURL *url.URL, transport http.RoundTripper, key string) Dialer { + return func(ctx context.Context) (DRPCClient, error) { + serveURL, err := serverURL.Parse("/api/v2/ai-gateway/serve") + if err != nil { + return nil, xerrors.Errorf("parse url: %w", err) + } + query := serveURL.Query() + query.Add(aibridgedproto.VersionQueryParam, aibridgedproto.CurrentVersion.String()) + serveURL.RawQuery = query.Encode() + + headers := http.Header{} + headers.Set(codersdk.BuildVersionHeader, buildinfo.Version()) + headers.Set(codersdk.AIGatewayKeyHeader, key) + + httpClient := &http.Client{ + Transport: transport, + } + // nolint:bodyclose // ReadBodyAsError closes the body; success path hands off to the websocket conn. + conn, res, err := websocket.Dial(ctx, serveURL.String(), &websocket.DialOptions{ + HTTPClient: httpClient, + CompressionMode: websocket.CompressionDisabled, + HTTPHeader: headers, + }) + if err != nil { + if res == nil { + return nil, err + } + return nil, readAIGatewayServeError(res) + } + config := yamux.DefaultConfig() + config.LogOutput = io.Discard + // Use a background context because the caller closes the client + // (and thus the multiplexed session) explicitly. + _, wsNetConn := codersdk.WebsocketNetConn(context.Background(), conn, websocket.MessageBinary) + conn.SetReadLimit(drpcsdk.YamuxDefaultStreamWindowSize) + session, err := yamux.Client(wsNetConn, config) + if err != nil { + _ = conn.Close(websocket.StatusGoingAway, "") + _ = wsNetConn.Close() + return nil, xerrors.Errorf("multiplex client: %w", err) + } + + dconn := drpcsdk.MultiplexedConn(session) + return &Client{ + Conn: dconn, + DRPCRecorderClient: aibridgedproto.NewDRPCRecorderClient(dconn), + DRPCMCPConfiguratorClient: aibridgedproto.NewDRPCMCPConfiguratorClient(dconn), + DRPCAuthorizerClient: aibridgedproto.NewDRPCAuthorizerClient(dconn), + DRPCProviderConfiguratorClient: aibridgedproto.NewDRPCProviderConfiguratorClient(dconn), + }, nil + } +} diff --git a/coderd/aibridged/http.go b/coderd/aibridged/http.go index 9927cb1b1dc07..7c9fef2cd5437 100644 --- a/coderd/aibridged/http.go +++ b/coderd/aibridged/http.go @@ -117,7 +117,7 @@ func (s *Server) ServeHTTP(rw http.ResponseWriter, r *http.Request) { r.Header.Del("X-Api-Key") } - client, err := s.Client() + client, err := s.ClientContext(ctx) if err != nil { logger.Warn(ctx, "failed to connect to coderd", slog.Error(err)) http.Error(rw, ErrConnect.Error(), http.StatusServiceUnavailable) diff --git a/coderd/aibridged/proto/version.go b/coderd/aibridged/proto/version.go index fea4d6d090c7e..bb484eeb1b73c 100644 --- a/coderd/aibridged/proto/version.go +++ b/coderd/aibridged/proto/version.go @@ -17,6 +17,11 @@ const ( CurrentMinor = 1 ) +// VersionQueryParam is the URL query parameter the standalone AI Gateway +// uses to advertise its aibridged API version when dialing coderd's serve +// endpoint, and that coderd reads to negotiate compatibility. +const VersionQueryParam = "version" + // CurrentVersion is the current aibridged API version. // Breaking changes to the aibridged API **MUST** increment CurrentMajor above. // Non-breaking changes to the aibridged API **MUST** increment CurrentMinor diff --git a/docs/reference/cli/ai-gateway.md b/docs/reference/cli/ai-gateway.md index 7666d36282732..1b8d2401d3c21 100644 --- a/docs/reference/cli/ai-gateway.md +++ b/docs/reference/cli/ai-gateway.md @@ -11,6 +11,7 @@ coder ai-gateway ## Subcommands -| Name | Purpose | -|-------------------------------------------|------------------------| -| [keys](./ai-gateway_keys.md) | Manage AI Gateway keys | +| Name | Purpose | +|---------------------------------------------|------------------------------------| +| [start](./ai-gateway_start.md) | Run a standalone AI Gateway server | +| [keys](./ai-gateway_keys.md) | Manage AI Gateway keys | diff --git a/docs/reference/cli/ai-gateway_start.md b/docs/reference/cli/ai-gateway_start.md new file mode 100644 index 0000000000000..f5c92b1e8ac5d --- /dev/null +++ b/docs/reference/cli/ai-gateway_start.md @@ -0,0 +1,141 @@ + +# ai-gateway start + +Run a standalone AI Gateway server + +## Usage + +```console +coder ai-gateway start [flags] +``` + +## Description + +```console +Runs a standalone replica of the AI Gateway. Standalone replicas serve LLM client traffic on a dedicated HTTP listener and connect to coderd using the Coder deployment URL and an AI Gateway key. + +Set --url or CODER_URL to the Coder deployment address, and set --key (CODER_AI_GATEWAY_KEY) or --key-file (CODER_AI_GATEWAY_KEY_FILE). A user login or session token is not required. +``` + +## Options + +### --key + +| | | +|-------------|------------------------------------| +| Type | string | +| Environment | $CODER_AI_GATEWAY_KEY | + +The AI Gateway key used to authenticate to coderd. + +### --key-file + +| | | +|-------------|-----------------------------------------| +| Type | string | +| Environment | $CODER_AI_GATEWAY_KEY_FILE | + +Path to a file containing the AI Gateway key used to authenticate to coderd. + +### --http-address + +| | | +|-------------|---------------------------------------------| +| Type | string | +| Environment | $CODER_AI_GATEWAY_HTTP_ADDRESS | +| Default | 127.0.0.1:4001 | + +The bind address to serve incoming AI Gateway client traffic. + +### --tls-cert-file + +| | | +|-------------|----------------------------------------------| +| Type | string | +| Environment | $CODER_AI_GATEWAY_TLS_CERT_FILE | + +Path to a PEM-encoded TLS certificate. Enables TLS termination when set together with --tls-key-file. + +### --tls-key-file + +| | | +|-------------|---------------------------------------------| +| Type | string | +| Environment | $CODER_AI_GATEWAY_TLS_KEY_FILE | + +Path to a PEM-encoded TLS private key. Enables TLS termination when set together with --tls-cert-file. + +### --verbose + +| | | +|-------------|----------------------------------------| +| Type | bool | +| Environment | $CODER_AI_GATEWAY_VERBOSE | +| Default | false | + +Output debug-level logs. + +### --ai-gateway-max-concurrency + +| | | +|-------------|------------------------------------------------| +| Type | int | +| Environment | $CODER_AI_GATEWAY_MAX_CONCURRENCY | +| YAML | ai_gateway.max_concurrency | +| Default | 0 | + +Maximum number of concurrent AI Gateway requests per replica. Set to 0 to disable (unlimited). + +### --ai-gateway-rate-limit + +| | | +|-------------|-------------------------------------------| +| Type | int | +| Environment | $CODER_AI_GATEWAY_RATE_LIMIT | +| YAML | ai_gateway.rate_limit | +| Default | 0 | + +Maximum number of AI Gateway requests per second per replica. Set to 0 to disable (unlimited). + +### --ai-gateway-send-actor-headers + +| | | +|-------------|---------------------------------------------------| +| Type | bool | +| Environment | $CODER_AI_GATEWAY_SEND_ACTOR_HEADERS | +| YAML | ai_gateway.send_actor_headers | +| Default | false | + +Once enabled, extra headers will be added to upstream requests to identify the user (actor) making requests to AI Gateway. This is only needed if you are using a proxy between AI Gateway and an upstream AI provider. This will send X-Ai-Bridge-Actor-Id (the ID of the user making the request) and X-Ai-Bridge-Actor-Metadata-Username (their username). + +### --ai-gateway-dump-dir + +| | | +|-------------|-----------------------------------------| +| Type | string | +| Environment | $CODER_AI_GATEWAY_DUMP_DIR | +| YAML | ai_gateway.api_dump_dir | + +Base directory for dumping AI Gateway request/response pairs to disk for debugging. When set, each provider writes under a subdirectory named after the provider. Sensitive headers are redacted. Leave empty to disable. + +### --ai-gateway-allow-byok + +| | | +|-------------|-------------------------------------------| +| Type | bool | +| Environment | $CODER_AI_GATEWAY_ALLOW_BYOK | +| YAML | ai_gateway.allow_byok | +| Default | true | + +Allow users to provide their own LLM API keys or subscriptions. When disabled, only centralized key authentication is permitted. + +### --ai-gateway-circuit-breaker-enabled + +| | | +|-------------|--------------------------------------------------------| +| Type | bool | +| Environment | $CODER_AI_GATEWAY_CIRCUIT_BREAKER_ENABLED | +| YAML | ai_gateway.circuit_breaker_enabled | +| Default | false | + +Enable the circuit breaker to protect against cascading failures from upstream AI provider overload (503, 529). diff --git a/enterprise/cli/aigateway.go b/enterprise/cli/aigateway.go index 2844ef43dc1b1..e1ac8c9add660 100644 --- a/enterprise/cli/aigateway.go +++ b/enterprise/cli/aigateway.go @@ -20,6 +20,7 @@ func (r *RootCmd) aiGateway() *serpent.Command { return inv.Command.HelpHandler(inv) }, Children: []*serpent.Command{ + r.aiGatewayStart(), r.aiGatewayKeys(), }, } diff --git a/enterprise/cli/aigatewaystart.go b/enterprise/cli/aigatewaystart.go new file mode 100644 index 0000000000000..cd84ef3ac969a --- /dev/null +++ b/enterprise/cli/aigatewaystart.go @@ -0,0 +1,309 @@ +//go:build !slim + +package cli + +import ( + "context" + "errors" + "net" + "net/http" + "os" + "strings" + "time" + + "github.com/prometheus/client_golang/prometheus" + tracenoop "go.opentelemetry.io/otel/trace/noop" + "golang.org/x/xerrors" + + "cdr.dev/slog/v3" + "cdr.dev/slog/v3/sloggers/sloghuman" + "github.com/coder/coder/v2/aibridge" + agpl "github.com/coder/coder/v2/cli" + "github.com/coder/coder/v2/coderd/aibridged" + "github.com/coder/coder/v2/codersdk" + "github.com/coder/coder/v2/enterprise/coderd" + "github.com/coder/retry" + "github.com/coder/serpent" +) + +const ( + shutdownTimeout = 5 * time.Minute + + keyFlagsExclusiveErr = "--key and --key-file options are mutually exclusive" + keyFlagsMissingErr = "an AI Gateway key is required, set --key (CODER_AI_GATEWAY_KEY) or --key-file (CODER_AI_GATEWAY_KEY_FILE)" +) + +// aiGatewayStart runs the AI Gateway as a standalone process. +func (r *RootCmd) aiGatewayStart() *serpent.Command { + var ( + key string + keyFile string + httpAddress string + tlsCertFile string + tlsKeyFile string + verbose bool + ) + + vals := new(codersdk.DeploymentValues) + + cmd := &serpent.Command{ + Use: "start", + Short: "Run a standalone AI Gateway server", + Long: "Runs a standalone replica of the AI Gateway. Standalone replicas " + + "serve LLM client traffic on a dedicated HTTP listener and connect " + + "to coderd using the Coder deployment URL and an AI Gateway key.\n\n" + + "Set --url or CODER_URL to the Coder deployment address, and set " + + "--key (CODER_AI_GATEWAY_KEY) or --key-file " + + "(CODER_AI_GATEWAY_KEY_FILE). A user login or session token is " + + "not required.", + Handler: func(inv *serpent.Invocation) error { + signalCtx, stop := inv.SignalNotifyContext(inv.Context(), agpl.StopSignals...) + defer stop() + + resolvedKey, err := resolveAIGatewayKey(key, keyFile) + if err != nil { + return err + } + + // TLS is opt-in and requires both files; setting only one is + // an error. Default is plain HTTP. + if (tlsCertFile == "") != (tlsKeyFile == "") { + return xerrors.New("--tls-cert-file and --tls-key-file options must be provided together") + } + + serverURL, transport, err := r.ResolveClientConnection() + if err != nil { + if errors.Is(err, agpl.ErrClientURLNotConfigured) { + return xerrors.New("AI Gateway requires --url or CODER_URL to point at the Coder deployment") + } + return xerrors.Errorf("configure Coder deployment connection: %w", err) + } + + logger := slog.Make(sloghuman.Sink(inv.Stderr)) + if verbose { + logger = logger.Leveled(slog.LevelDebug) + } + + // Metrics and tracing are not exposed by standalone mode yet + // (TODO AIGOV-317), but the pool and the reloader require a metrics + // object and a tracer. + registry := prometheus.NewRegistry() + metrics := aibridge.NewMetrics(registry) + providerMetrics := aibridged.NewMetrics(registry) + tracer := tracenoop.NewTracerProvider().Tracer("aibridged") + + // Standalone Gateway starts with an empty pool. Providers are + // fetched later via GetAIProviders DRPC and pool is updated. + pool, err := aibridged.NewCachedBridgePool(aibridged.DefaultPoolOptions, nil, logger.Named("pool"), metrics, tracer) + if err != nil { + return xerrors.Errorf("create request pool: %w", err) + } + + dialer := aibridged.NewWebsocketDialer(serverURL, transport, resolvedKey) + aibridgedCtx, aibridgedCancel := context.WithCancel(context.Background()) + defer aibridgedCancel() + srv, err := aibridged.New(aibridgedCtx, pool, dialer, logger.Named("aibridged"), tracer) + if err != nil { + return xerrors.Errorf("start AI Gateway daemon: %w", err) + } + defer srv.Close() + + // Fetch the initial provider set from coderd, retrying until + // success. + // TODO(AIGOV-465): the standalone gateway has no refresh trigger + // yet, so this runs once on startup. + clientFn := func() (aibridged.DRPCClient, error) { + return srv.ClientContext(signalCtx) + } + providerLogger := logger.Named("aibridge.providers") + reloader := agpl.NewPoolRPCReloader(pool, clientFn, vals.AI.BridgeConfig, providerLogger, metrics, providerMetrics) + if err := loadProviders(signalCtx, reloader, providerLogger, srv.Done()); err != nil { + if signalCtx.Err() != nil { + logger.Info(signalCtx, "shutting down standalone AI Gateway") + return nil + } + return xerrors.Errorf("initialize ai providers: %w", err) + } + + mw := coderd.AIGatewayDataPlaneMiddleware(vals.AI.BridgeConfig) + + // The standalone listener is dedicated to Gateway traffic, so + // the daemon is served at the root. The /api/v2/ai-gateway + // and /api/v2/aibridge/ aliases are added for compatibility + // with the embedded route. + mux := http.NewServeMux() + mux.Handle("/api/v2/aibridge/", mw(http.StripPrefix("/api/v2/aibridge", srv))) + mux.Handle("/api/v2/ai-gateway/", mw(http.StripPrefix("/api/v2/ai-gateway", srv))) + mux.Handle("/", mw(srv)) + + listener, err := net.Listen("tcp", httpAddress) + if err != nil { + return xerrors.Errorf("listen on %q: %w", httpAddress, err) + } + defer listener.Close() + + logger.Info(signalCtx, "standalone AI Gateway listening", + slog.F("address", listener.Addr().String()), + slog.F("coder_url", serverURL.String()), + slog.F("tls", tlsCertFile != ""), + ) + + httpServer := &http.Server{ + Handler: mux, + ReadHeaderTimeout: time.Minute, + } + + serveErr := make(chan error, 1) + go func() { + if tlsCertFile != "" { + serveErr <- httpServer.ServeTLS(listener, tlsCertFile, tlsKeyFile) + } else { + serveErr <- httpServer.Serve(listener) + } + }() + + var aibridgedErr error + select { + case <-signalCtx.Done(): + logger.Info(signalCtx, "shutting down standalone AI Gateway") + case <-srv.Done(): + aibridgedErr = srv.Err() + case err := <-serveErr: + if err != nil && !errors.Is(err, http.ErrServerClosed) { + return xerrors.Errorf("serve: %w", err) + } + } + + shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), shutdownTimeout) + defer shutdownCancel() + if err := httpServer.Shutdown(shutdownCtx); err != nil { + return xerrors.Errorf("shutdown http server: %w", err) + } + if aibridgedErr != nil { + return xerrors.Errorf("AI Gateway daemon exited: %w", aibridgedErr) + } + return nil + }, + } + + cmd.Options = serpent.OptionSet{ + { + Flag: "key", + Env: "CODER_AI_GATEWAY_KEY", + Description: "The AI Gateway key used to authenticate to coderd.", + Value: serpent.StringOf(&key), + }, + { + Flag: "key-file", + Env: "CODER_AI_GATEWAY_KEY_FILE", + Description: "Path to a file containing the AI Gateway key used to authenticate to coderd.", + Value: serpent.StringOf(&keyFile), + }, + { + Flag: "http-address", + Env: "CODER_AI_GATEWAY_HTTP_ADDRESS", + Description: "The bind address to serve incoming AI Gateway client traffic.", + Default: "127.0.0.1:4001", + Value: serpent.StringOf(&httpAddress), + }, + { + Flag: "tls-cert-file", + Env: "CODER_AI_GATEWAY_TLS_CERT_FILE", + Description: "Path to a PEM-encoded TLS certificate. Enables TLS termination when set together with --tls-key-file.", + Value: serpent.StringOf(&tlsCertFile), + }, + { + Flag: "tls-key-file", + Env: "CODER_AI_GATEWAY_TLS_KEY_FILE", + Description: "Path to a PEM-encoded TLS private key. Enables TLS termination when set together with --tls-cert-file.", + Value: serpent.StringOf(&tlsKeyFile), + }, + { + Flag: "verbose", + Env: "CODER_AI_GATEWAY_VERBOSE", + Description: "Output debug-level logs.", + Value: serpent.BoolOf(&verbose), + Default: "false", + }, + } + + // Standalone Gateway only uses part of the options from "AI Gateway" group. + // Other options from the group are coderd-only (eg. budget, provider-seeding). + standaloneOpts := map[string]struct{}{ + "CODER_AI_GATEWAY_ALLOW_BYOK": {}, + "CODER_AI_GATEWAY_SEND_ACTOR_HEADERS": {}, + "CODER_AI_GATEWAY_DUMP_DIR": {}, + "CODER_AI_GATEWAY_CIRCUIT_BREAKER_ENABLED": {}, + "CODER_AI_GATEWAY_CIRCUIT_BREAKER_FAILURE_THRESHOLD": {}, + "CODER_AI_GATEWAY_CIRCUIT_BREAKER_INTERVAL": {}, + "CODER_AI_GATEWAY_CIRCUIT_BREAKER_TIMEOUT": {}, + "CODER_AI_GATEWAY_CIRCUIT_BREAKER_MAX_REQUESTS": {}, + "CODER_AI_GATEWAY_MAX_CONCURRENCY": {}, + "CODER_AI_GATEWAY_RATE_LIMIT": {}, + } + + var aiGatewayOpts serpent.OptionSet + for _, opt := range vals.Options() { + if opt.Group == nil || opt.Group.Name != "AI Gateway" { + continue + } + if _, ok := standaloneOpts[opt.Env]; !ok { + continue + } + aiGatewayOpts = append(aiGatewayOpts, opt) + } + + cmd.Options = append(cmd.Options, aiGatewayOpts...) + + return cmd +} + +// resolveAIGatewayKey resolves key from --key or --key-file flags. +// If both are set, an error is returned. If neither is set, an empty string is returned. +func resolveAIGatewayKey(key string, keyFile string) (string, error) { + if key != "" && keyFile != "" { + return "", xerrors.New(keyFlagsExclusiveErr) + } + if key == "" && keyFile == "" { + return "", xerrors.New(keyFlagsMissingErr) + } + if keyFile == "" { + return key, nil + } + data, err := os.ReadFile(keyFile) + if err != nil { + return "", xerrors.Errorf("read AI Gateway key file %q: %w", keyFile, err) + } + return strings.TrimSpace(string(data)), nil +} + +// loadProviders performs the standalone gateway's initial provider +// load by driving reloader until it succeeds or ctx is canceled. The reloader +// owns the actual fetch/build/replace/metrics work; the reloader's underlying +// client blocks until the daemon connects to coderd, and the fetch may still +// fail transiently (e.g. mid-seed contention or a dropped connection), so the +// reload is retried with backoff. A successful empty provider list is a valid +// result and ends the loop. +// +// TODO(AIGOV-465): the standalone gateway has no provider-change refresh +// trigger yet, so this runs once on startup; provider add/enable will not +// propagate to a running standalone gateway. +func loadProviders(ctx context.Context, reloader aibridged.ProviderReloader, logger slog.Logger, aibridgedDone <-chan struct{}) error { + for r := retry.New(50*time.Millisecond, 10*time.Second); r.Wait(ctx); { + if err := reloader.Reload(ctx); err != nil { + select { + case <-aibridgedDone: + return err + default: + } + logger.Warn(ctx, "failed to load ai providers, will retry", slog.Error(err)) + continue + } + logger.Info(ctx, "loaded ai providers from coderd") + return nil + } + if cause := context.Cause(ctx); cause != nil { + return cause + } + return ctx.Err() +} diff --git a/enterprise/cli/aigatewaystart_internal_test.go b/enterprise/cli/aigatewaystart_internal_test.go new file mode 100644 index 0000000000000..db5309bdf4f7e --- /dev/null +++ b/enterprise/cli/aigatewaystart_internal_test.go @@ -0,0 +1,217 @@ +//go:build !slim + +package cli + +import ( + "context" + "os" + "path/filepath" + "sync/atomic" + "testing" + + "github.com/stretchr/testify/require" + "golang.org/x/xerrors" + + "cdr.dev/slog/v3" + "github.com/coder/coder/v2/testutil" +) + +// blockingReloader blocks in Reload until the context is canceled, then +// returns its error. It models the standalone gateway's initial reload +// waiting on a daemon connection to an unreachable coderd. +type blockingReloader struct { + started chan struct{} +} + +func (r *blockingReloader) Reload(ctx context.Context) error { + select { + case r.started <- struct{}{}: + default: + } + <-ctx.Done() + return ctx.Err() +} + +// failThenSucceedReloader fails the first failUntil reloads, then succeeds, +// modeling a coderd connection or provider fetch that recovers after a few +// transient failures. +type failThenSucceedReloader struct { + calls atomic.Int32 + failUntil int32 +} + +func (r *failThenSucceedReloader) Reload(_ context.Context) error { + if r.calls.Add(1) <= r.failUntil { + return xerrors.New("transient failure") + } + return nil +} + +// alwaysFailReloader returns the same error every time Reload is called. +type alwaysFailReloader struct { + calls atomic.Int32 + err error + after func() + called chan struct{} +} + +func (r *alwaysFailReloader) Reload(context.Context) error { + r.calls.Add(1) + if r.after != nil { + r.after() + } + select { + case r.called <- struct{}{}: + default: + } + return r.err +} + +// TestLoadProviders_Interruptible verifies that a stop signal, +// modeled by canceling the context, unblocks the initial provider load even +// when the reloader is stuck waiting for coderd. This guards the standalone +// "ai-gateway start" command against the regression where startup could not +// be interrupted. +func TestLoadProviders_Interruptible(t *testing.T) { + t.Parallel() + + // testCtx bounds the test and drives the channel receives; runCtx is the + // context handed to loadProviders and is canceled to model a + // stop signal. They are distinct so the receives still work after the + // signal context is canceled. + testCtx := testutil.Context(t, testutil.WaitShort) + runCtx, cancel := context.WithCancel(testCtx) + defer cancel() + + reloader := &blockingReloader{started: make(chan struct{}, 1)} + logger := slog.Make() + + done := make(chan error, 1) + go func() { + done <- loadProviders(runCtx, reloader, logger, nil) + }() + + // Wait for the reload to be in-flight, then cancel as a signal would. + testutil.RequireReceive(testCtx, t, reloader.started) + cancel() + + err := testutil.RequireReceive(testCtx, t, done) + require.ErrorIs(t, err, context.Canceled) +} + +// TestLoadProviders_RetrySucceeds verifies loadProviders keeps retrying past +// transient failures and returns nil once a reload succeeds. This guards the +// retry contract: replacing the loop's continue with a return would fail here. +func TestLoadProviders_RetrySucceeds(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitShort) + reloader := &failThenSucceedReloader{failUntil: 2} + + require.NoError(t, loadProviders(ctx, reloader, slog.Make(), nil)) + require.GreaterOrEqual(t, reloader.calls.Load(), int32(3)) +} + +func TestLoadProviders_AIBridgedDoneStopsRetry(t *testing.T) { + t.Parallel() + + errMsg := "aibridged fatal" + ctx := testutil.Context(t, testutil.WaitShort) + aibridgedDone := make(chan struct{}) + reloader := &alwaysFailReloader{ + err: xerrors.New(errMsg), + called: make(chan struct{}, 1), + after: func() { + close(aibridgedDone) + }, + } + + err := loadProviders(ctx, reloader, slog.Make(), aibridgedDone) + require.ErrorContains(t, err, errMsg) + require.Equal(t, int32(1), reloader.calls.Load()) +} + +func TestResolveAIGatewayKey(t *testing.T) { + t.Parallel() + + keyFile := filepath.Join(t.TempDir(), "gateway.key") + require.NoError(t, os.WriteFile(keyFile, []byte("file-key\n"), 0o600)) + + tests := []struct { + name string + key string + keyFile string + want string + wantErr string + }{ + { + name: "Nothing set", + wantErr: keyFlagsMissingErr, + }, + { + name: "Key", + key: "flag-key", + want: "flag-key", + }, + { + name: "KeyFile", + keyFile: keyFile, + want: "file-key", + }, + { + name: "MutuallyExclusive", + key: "flag-key", + keyFile: keyFile, + wantErr: keyFlagsExclusiveErr, + }, + { + name: "MissingKeyFile", + keyFile: filepath.Join(t.TempDir(), "missing.key"), + wantErr: "read AI Gateway key file", + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + got, err := resolveAIGatewayKey(tc.key, tc.keyFile) + if tc.wantErr != "" { + require.ErrorContains(t, err, tc.wantErr) + return + } + require.NoError(t, err) + require.Equal(t, tc.want, got) + }) + } +} + +func TestAIGatewayStart_DeploymentOptions(t *testing.T) { + t.Parallel() + + cmd := (&RootCmd{}).aiGatewayStart() + + // Standalone Gateway only consumes deployment options used in LLM traffic. + // Coderd-only settings such as provider seeds, retention, + // structured logging, and Coder MCP injection must stay server-only. + var got []string + for _, opt := range cmd.Options { + if opt.Group != nil && opt.Group.Name == "AI Gateway" { + got = append(got, opt.Env) + } + } + + want := []string{ + "CODER_AI_GATEWAY_ALLOW_BYOK", + "CODER_AI_GATEWAY_CIRCUIT_BREAKER_ENABLED", + "CODER_AI_GATEWAY_CIRCUIT_BREAKER_FAILURE_THRESHOLD", + "CODER_AI_GATEWAY_CIRCUIT_BREAKER_INTERVAL", + "CODER_AI_GATEWAY_CIRCUIT_BREAKER_MAX_REQUESTS", + "CODER_AI_GATEWAY_CIRCUIT_BREAKER_TIMEOUT", + "CODER_AI_GATEWAY_DUMP_DIR", + "CODER_AI_GATEWAY_MAX_CONCURRENCY", + "CODER_AI_GATEWAY_RATE_LIMIT", + "CODER_AI_GATEWAY_SEND_ACTOR_HEADERS", + } + require.ElementsMatch(t, want, got) +} diff --git a/enterprise/cli/aigatewaystart_slim.go b/enterprise/cli/aigatewaystart_slim.go new file mode 100644 index 0000000000000..c43f553e83c4a --- /dev/null +++ b/enterprise/cli/aigatewaystart_slim.go @@ -0,0 +1,24 @@ +//go:build slim + +package cli + +import ( + agplcli "github.com/coder/coder/v2/cli" + "github.com/coder/serpent" +) + +func (r *RootCmd) aiGatewayStart() *serpent.Command { + cmd := &serpent.Command{ + Use: "start", + Short: "Run a standalone AI Gateway server", + // We accept RawArgs so all commands and flags are accepted. + RawArgs: true, + Hidden: true, + Handler: func(inv *serpent.Invocation) error { + agplcli.SlimUnsupported(inv.Stderr, "ai-gateway start") + return nil + }, + } + + return cmd +} diff --git a/enterprise/cli/testdata/coder_ai-gateway_--help.golden b/enterprise/cli/testdata/coder_ai-gateway_--help.golden index 2c0d35058b767..7569e8e89f8ff 100644 --- a/enterprise/cli/testdata/coder_ai-gateway_--help.golden +++ b/enterprise/cli/testdata/coder_ai-gateway_--help.golden @@ -6,7 +6,8 @@ USAGE: Manage AI Gateway SUBCOMMANDS: - keys Manage AI Gateway keys + keys Manage AI Gateway keys + start Run a standalone AI Gateway server ——— Run `coder --help` for a list of global options. diff --git a/enterprise/cli/testdata/coder_ai-gateway_start_--help.golden b/enterprise/cli/testdata/coder_ai-gateway_start_--help.golden new file mode 100644 index 0000000000000..8156fbbf122a7 --- /dev/null +++ b/enterprise/cli/testdata/coder_ai-gateway_start_--help.golden @@ -0,0 +1,70 @@ +coder v0.0.0-devel + +USAGE: + coder ai-gateway start [flags] + + Run a standalone AI Gateway server + + Runs a standalone replica of the AI Gateway. Standalone replicas serve LLM + client traffic on a dedicated HTTP listener and connect to coderd using the + Coder deployment URL and an AI Gateway key. + + Set --url or CODER_URL to the Coder deployment address, and set --key + (CODER_AI_GATEWAY_KEY) or --key-file (CODER_AI_GATEWAY_KEY_FILE). A user login + or session token is not required. + +OPTIONS: + --http-address string, $CODER_AI_GATEWAY_HTTP_ADDRESS (default: 127.0.0.1:4001) + The bind address to serve incoming AI Gateway client traffic. + + --key string, $CODER_AI_GATEWAY_KEY + The AI Gateway key used to authenticate to coderd. + + --key-file string, $CODER_AI_GATEWAY_KEY_FILE + Path to a file containing the AI Gateway key used to authenticate to + coderd. + + --tls-cert-file string, $CODER_AI_GATEWAY_TLS_CERT_FILE + Path to a PEM-encoded TLS certificate. Enables TLS termination when + set together with --tls-key-file. + + --tls-key-file string, $CODER_AI_GATEWAY_TLS_KEY_FILE + Path to a PEM-encoded TLS private key. Enables TLS termination when + set together with --tls-cert-file. + + --verbose bool, $CODER_AI_GATEWAY_VERBOSE (default: false) + Output debug-level logs. + +AI GATEWAY OPTIONS: + --ai-gateway-dump-dir string, $CODER_AI_GATEWAY_DUMP_DIR + Base directory for dumping AI Gateway request/response pairs to disk + for debugging. When set, each provider writes under a subdirectory + named after the provider. Sensitive headers are redacted. Leave empty + to disable. + + --ai-gateway-allow-byok bool, $CODER_AI_GATEWAY_ALLOW_BYOK (default: true) + Allow users to provide their own LLM API keys or subscriptions. When + disabled, only centralized key authentication is permitted. + + --ai-gateway-circuit-breaker-enabled bool, $CODER_AI_GATEWAY_CIRCUIT_BREAKER_ENABLED (default: false) + Enable the circuit breaker to protect against cascading failures from + upstream AI provider overload (503, 529). + + --ai-gateway-max-concurrency int, $CODER_AI_GATEWAY_MAX_CONCURRENCY (default: 0) + Maximum number of concurrent AI Gateway requests per replica. Set to 0 + to disable (unlimited). + + --ai-gateway-rate-limit int, $CODER_AI_GATEWAY_RATE_LIMIT (default: 0) + Maximum number of AI Gateway requests per second per replica. Set to 0 + to disable (unlimited). + + --ai-gateway-send-actor-headers bool, $CODER_AI_GATEWAY_SEND_ACTOR_HEADERS (default: false) + Once enabled, extra headers will be added to upstream requests to + identify the user (actor) making requests to AI Gateway. This is only + needed if you are using a proxy between AI Gateway and an upstream AI + provider. This will send X-Ai-Bridge-Actor-Id (the ID of the user + making the request) and X-Ai-Bridge-Actor-Metadata-Username (their + username). + +——— +Run `coder --help` for a list of global options. diff --git a/enterprise/coderd/aibridge.go b/enterprise/coderd/aibridge.go index c454bc4835be9..d4f8998252820 100644 --- a/enterprise/coderd/aibridge.go +++ b/enterprise/coderd/aibridge.go @@ -72,12 +72,6 @@ func aiGatewayHTTPHandler(api *API, middlewares ...func(http.Handler) http.Handl // under /aibridge. The stripPrefix parameter selects which URL prefix // to strip before forwarding to the in-memory aibridged handler. func aiBridgeRoutes(api *API, stripPrefix string, middlewares ...func(http.Handler) http.Handler) func(r chi.Router) { - // Build the overload protection middleware chain for the aibridged handler. - // These limits are applied per-replica. - bridgeCfg := api.DeploymentValues.AI.BridgeConfig - concurrencyLimiter := httpmw.ConcurrencyLimit(bridgeCfg.MaxConcurrency.Value(), "AI Gateway") - rateLimiter := httpmw.RateLimitByAuthToken(int(bridgeCfg.RateLimit.Value()), aiBridgeRateLimitWindow) - return func(r chi.Router) { r.Use(api.RequireFeatureMW(codersdk.FeatureAIBridge)) r.Group(func(r chi.Router) { @@ -88,10 +82,10 @@ func aiBridgeRoutes(api *API, stripPrefix string, middlewares ...func(http.Handl r.Get("/clients", api.aiBridgeListClients) }) - // Apply overload protection middleware to the aibridged handler. - // Concurrency limit is checked first for faster rejection under load. + // Apply the shared per-request data-plane middleware (per-replica + // overload protection plus BYOK gating) to the aibridged handler. r.Group(func(r chi.Router) { - r.Use(concurrencyLimiter, rateLimiter) + r.Use(AIGatewayDataPlaneMiddleware(api.DeploymentValues.AI.BridgeConfig)) // This is a bit funky but since aibridge only exposes a HTTP // handler, this is how it has to be. r.HandleFunc("/*", func(rw http.ResponseWriter, r *http.Request) { @@ -103,16 +97,6 @@ func aiBridgeRoutes(api *API, stripPrefix string, middlewares ...func(http.Handl return } - // Reject BYOK requests when the deployment has not - // enabled bring-your-own-key mode. - if agplaibridge.IsBYOK(r.Header) && !bridgeCfg.AllowBYOK.Value() { - httpapi.Write(r.Context(), rw, http.StatusForbidden, codersdk.Response{ - Message: "Bring Your Own Key (BYOK) mode is not enabled.", - Detail: "Contact your administrator to enable it with --aibridge-allow-byok.", - }) - return - } - // Strip the prefix and relay to the aibridged handler. http.StripPrefix(stripPrefix, handler).ServeHTTP(rw, r) }) @@ -120,6 +104,33 @@ func aiBridgeRoutes(api *API, stripPrefix string, middlewares ...func(http.Handl } } +// AIGatewayDataPlaneMiddleware returns the per-request middleware chain that +// guards the AI Gateway data-plane handler. It is the single source of truth +// shared by the embedded route and the standalone gateway. +func AIGatewayDataPlaneMiddleware(cfg codersdk.AIBridgeConfig) func(http.Handler) http.Handler { + concurrencyLimiter := httpmw.ConcurrencyLimit(cfg.MaxConcurrency.Value(), "AI Gateway") + rateLimiter := httpmw.RateLimitByAuthToken(int(cfg.RateLimit.Value()), aiBridgeRateLimitWindow) + byokGuard := aiGatewayBYOKGuard(cfg) + return func(next http.Handler) http.Handler { + return concurrencyLimiter(rateLimiter(byokGuard(next))) + } +} + +func aiGatewayBYOKGuard(cfg codersdk.AIBridgeConfig) func(http.Handler) http.Handler { + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(rw http.ResponseWriter, r *http.Request) { + if agplaibridge.IsBYOK(r.Header) && !cfg.AllowBYOK.Value() { + httpapi.Write(r.Context(), rw, http.StatusForbidden, codersdk.Response{ + Message: "Bring Your Own Key (BYOK) mode is not enabled.", + Detail: "Contact your administrator to enable it with --ai-gateway-allow-byok.", + }) + return + } + next.ServeHTTP(rw, r) + }) + } +} + // aiBridgeListSessions returns AI Bridge sessions (aggregated interceptions). // // @Summary List AI Gateway sessions diff --git a/enterprise/coderd/aibridgeserve.go b/enterprise/coderd/aibridgeserve.go index f2f9fcc54dbf2..7cd8586a0d7e8 100644 --- a/enterprise/coderd/aibridgeserve.go +++ b/enterprise/coderd/aibridgeserve.go @@ -66,7 +66,7 @@ func (api *API) aiGatewayServe(rw http.ResponseWriter, r *http.Request) { return } - clientAPIVersion := r.URL.Query().Get("version") + clientAPIVersion := r.URL.Query().Get(aibridgedproto.VersionQueryParam) clientCoderVersion := r.Header.Get(codersdk.BuildVersionHeader) logger := api.Logger.Named("aigateway-serve").With( slog.F("remote_addr", r.RemoteAddr), @@ -88,7 +88,7 @@ func (api *API) aiGatewayServe(rw http.ResponseWriter, r *http.Request) { httpapi.Write(keyCtx, rw, http.StatusBadRequest, codersdk.Response{ Message: "Incompatible or unparsable version", Validations: []codersdk.ValidationError{ - {Field: "version", Detail: err.Error()}, + {Field: aibridgedproto.VersionQueryParam, Detail: err.Error()}, {Field: "client_api_version", Detail: clientAPIVersion}, {Field: "server_api_version", Detail: aibridgedproto.CurrentVersion.String()}, }, @@ -131,7 +131,7 @@ func (api *API) aiGatewayServe(rw http.ResponseWriter, r *http.Request) { if _, err := aiGatewayUpdateKeyLastHeartbeat(connCtx, api, gatewayKey.ID); err != nil { logger.Warn(connCtx, "update ai gateway key last heartbeat", slog.Error(err)) } - go aiGatewayTrackKeyUsage(connCtx, keyCtxCancel, api, gatewayKey.ID, logger) + go aiGatewayCheckEntitlementAndTrackKeyUsage(connCtx, keyCtxCancel, api, gatewayKey.ID, logger) mux := drpcmux.New() srv, err := aibridgedserver.NewServer( @@ -194,8 +194,11 @@ func aiGatewayUpdateKeyLastHeartbeat(ctx context.Context, api *API, keyID uuid.U return rows > 0, nil } -// aiGatewayTrackKeyUsage refreshes last_heartbeat_at for keyID on a fixed interval until ctx is canceled. -func aiGatewayTrackKeyUsage(ctx context.Context, ctxCancel context.CancelFunc, api *API, keyID uuid.UUID, logger slog.Logger) { +// aiGatewayCheckEntitlementAndTrackKeyUsage until ctx is canceled on a fixed interval: +// - refreshes last_heartbeat_at for keyID. +// - checks if key still exists, cancels ctx if it does not. +// - checks if the AI Gov entitlement is still enabled, cancels ctx if it is not. +func aiGatewayCheckEntitlementAndTrackKeyUsage(ctx context.Context, ctxCancel context.CancelFunc, api *API, keyID uuid.UUID, logger slog.Logger) { ticker, done := api.NewTicker(aiGatewayKeyHeartbeatInterval) defer done() @@ -214,6 +217,13 @@ func aiGatewayTrackKeyUsage(ctx context.Context, ctxCancel context.CancelFunc, a return } + // Close connection when the entitlement is revoked. + if !api.Entitlements.Enabled(codersdk.FeatureAIBridge) { + logger.Info(ctx, "ai gateway entitlement no longer enabled, closing connection") + ctxCancel() + return + } + if err != nil { if xerrors.Is(err, context.Canceled) { return diff --git a/enterprise/coderd/aibridgeserve_test.go b/enterprise/coderd/aibridgeserve_test.go index 76f3ffe178a1d..9a702f3af54aa 100644 --- a/enterprise/coderd/aibridgeserve_test.go +++ b/enterprise/coderd/aibridgeserve_test.go @@ -2,70 +2,63 @@ package coderd_test import ( "context" - "io" "net/http" "testing" "time" - "github.com/hashicorp/yamux" + "github.com/google/uuid" "github.com/stretchr/testify/require" + "google.golang.org/protobuf/types/known/timestamppb" + "github.com/coder/coder/v2/coderd/aibridged" aibridgedproto "github.com/coder/coder/v2/coderd/aibridged/proto" "github.com/coder/coder/v2/coderd/coderdtest" "github.com/coder/coder/v2/codersdk" - "github.com/coder/coder/v2/codersdk/drpcsdk" + entcoderd "github.com/coder/coder/v2/enterprise/coderd" "github.com/coder/coder/v2/enterprise/coderd/coderdenttest" "github.com/coder/coder/v2/enterprise/coderd/license" "github.com/coder/coder/v2/testutil" "github.com/coder/serpent" - "github.com/coder/websocket" ) -// dialAIGatewayServe dials /api/v2/ai-gateway/serve, authenticating with the given -// gateway key and API version. On a successful WebSocket upgrade it returns a -// yamux session and http.StatusSwitchingProtocols. Otherwise it returns a nil -// session and the HTTP status code coderd responded with. -func dialAIGatewayServe(ctx context.Context, t *testing.T, client *codersdk.Client, key string, version string) (*yamux.Session, int) { - t.Helper() +type versionOverridingRoundTripper struct { + baseTransport http.RoundTripper + overrideAPIVersion string +} - serverURL, err := client.URL.Parse("/api/v2/ai-gateway/serve") - require.NoError(t, err) - query := serverURL.Query() - if version != "" { - query.Set("version", version) +func (f versionOverridingRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) { + query := req.URL.Query() + query.Del(aibridgedproto.VersionQueryParam) + if f.overrideAPIVersion != "" { + query.Set(aibridgedproto.VersionQueryParam, f.overrideAPIVersion) } - serverURL.RawQuery = query.Encode() + req.URL.RawQuery = query.Encode() + return f.baseTransport.RoundTrip(req) +} + +func dialAIGatewayServe(ctx context.Context, t *testing.T, client *codersdk.Client, key string) (aibridged.DRPCClient, error) { + return dialAIGatewayServeWithVersion(ctx, t, client, key, nil) +} - headers := http.Header{} - if key != "" { - headers.Set(codersdk.AIGatewayKeyHeader, key) +func dialAIGatewayServeWithVersion(ctx context.Context, t *testing.T, client *codersdk.Client, key string, version *string) (aibridged.DRPCClient, error) { + t.Helper() + + transport := client.HTTPClient.Transport + if version != nil { + transport = versionOverridingRoundTripper{ + baseTransport: transport, + overrideAPIVersion: *version, + } } - conn, res, err := websocket.Dial(ctx, serverURL.String(), &websocket.DialOptions{ - HTTPClient: &http.Client{Transport: client.HTTPClient.Transport}, - CompressionMode: websocket.CompressionDisabled, - HTTPHeader: headers, - }) + dc, err := aibridged.NewWebsocketDialer(client.URL, transport, key)(ctx) if err != nil { - statusCode := 0 - if res != nil { - statusCode = res.StatusCode - _ = res.Body.Close() - } - return nil, statusCode + return nil, err } - cfg := yamux.DefaultConfig() - cfg.LogOutput = io.Discard - _, wsNetConn := codersdk.WebsocketNetConn(context.Background(), conn, websocket.MessageBinary) - conn.SetReadLimit(drpcsdk.YamuxDefaultStreamWindowSize) - session, err := yamux.Client(wsNetConn, cfg) - require.NoError(t, err) t.Cleanup(func() { - _ = session.Close() - _ = wsNetConn.Close() - _ = conn.Close(websocket.StatusNormalClosure, "") + _ = dc.DRPCConn().Close() }) - return session, http.StatusSwitchingProtocols + return dc, nil } func TestAIGatewayServeSuccess(t *testing.T) { @@ -78,20 +71,38 @@ func TestAIGatewayServeSuccess(t *testing.T) { created, err := client.CreateAIGatewayKey(ctx, codersdk.CreateAIGatewayKeyRequest{Name: "serve-success"}) require.NoError(t, err) - session, status := dialAIGatewayServe(ctx, t, client, created.Key, aibridgedproto.CurrentVersion.String()) - require.Equal(t, http.StatusSwitchingProtocols, status) - require.NotNil(t, session) + // Use NewWebsocketDialer that production code of standalone gateway uses + dc, err := dialAIGatewayServe(ctx, t, client, created.Key) + require.NoError(t, err) - // The Authorizer service should be served and authorize the owner's - // session token, exercising a full DRPC round trip over the WebSocket. - authorizer := aibridgedproto.NewDRPCAuthorizerClient(drpcsdk.MultiplexedConn(session)) - resp, err := authorizer.IsAuthorized(ctx, &aibridgedproto.IsAuthorizedRequest{ - Key: client.SessionToken(), - }) + // Exercise one RPC from each service in the DRPCClient union to verify the + // dialer wires every service and the serve mux registers them all. + + // DRPCAuthorizerClient + resp, err := dc.IsAuthorized(ctx, &aibridgedproto.IsAuthorizedRequest{Key: client.SessionToken()}) require.NoError(t, err) require.Equal(t, firstUser.UserID.String(), resp.GetOwnerId()) - // The session records liveness for the authenticating key. + // DRPCProviderConfiguratorClient + _, err = dc.GetAIProviders(ctx, &aibridgedproto.GetAIProvidersRequest{}) + require.NoError(t, err) + + // DRPCMCPConfiguratorClient + _, err = dc.GetMCPServerConfigs(ctx, &aibridgedproto.GetMCPServerConfigsRequest{UserId: firstUser.UserID.String()}) + require.NoError(t, err) + + // DRPCRecorderClient + _, err = dc.RecordInterception(ctx, &aibridgedproto.RecordInterceptionRequest{ + Id: uuid.NewString(), + InitiatorId: firstUser.UserID.String(), + ApiKeyId: "serve-success-key", + Provider: "openai", + Model: "gpt-4", + StartedAt: timestamppb.Now(), + }) + require.NoError(t, err) + + // Verify the session records liveness for the authenticating key. require.Eventually(t, func() bool { //nolint:gocritic // Owner role is needed for gateway key management. keys, err := client.ListAIGatewayKeys(ctx) @@ -124,48 +135,65 @@ func TestAIGatewayServeKeyAndVersionValidationErr(t *testing.T) { require.NoError(t, client.DeleteAIGatewayKey(ctx, revoked.ID)) tests := []struct { - name string - key string - version string - wantStatus int + name string + key string + version string + wantStatus int + wantMessage string + forbidErrMessage string }{ { - name: "MissingKey", - key: "", - version: aibridgedproto.CurrentVersion.String(), - wantStatus: http.StatusUnauthorized, + name: "MissingKey", + key: "", + version: aibridgedproto.CurrentVersion.String(), + wantStatus: http.StatusUnauthorized, + wantMessage: "AI Gateway key required.", + forbidErrMessage: "Try logging in", }, { - name: "InvalidKey", - key: "not-a-real-key", - version: aibridgedproto.CurrentVersion.String(), - wantStatus: http.StatusUnauthorized, + name: "InvalidKey", + key: "not-a-real-key", + version: aibridgedproto.CurrentVersion.String(), + wantStatus: http.StatusUnauthorized, + wantMessage: "AI Gateway key invalid.", + forbidErrMessage: "Try logging in", }, { - name: "RevokedKey", - key: revoked.Key, - version: aibridgedproto.CurrentVersion.String(), - wantStatus: http.StatusUnauthorized, + name: "RevokedKey", + key: revoked.Key, + version: aibridgedproto.CurrentVersion.String(), + wantStatus: http.StatusUnauthorized, + wantMessage: "AI Gateway key invalid.", + forbidErrMessage: "Try logging in", }, { - name: "IncompatibleVersion", - key: validKey, - version: "999.0", - wantStatus: http.StatusBadRequest, + name: "IncompatibleVersion", + key: validKey, + version: "999.0", + wantStatus: http.StatusBadRequest, + wantMessage: "Incompatible or unparsable version", }, { - name: "MissingVersion", - key: validKey, - version: "", - wantStatus: http.StatusBadRequest, + name: "MissingVersion", + key: validKey, + version: "", + wantStatus: http.StatusBadRequest, + wantMessage: "Incompatible or unparsable version", }, } for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { t.Parallel() - _, status := dialAIGatewayServe(t.Context(), t, client, tc.key, tc.version) - require.Equal(t, tc.wantStatus, status) + + _, err := dialAIGatewayServeWithVersion(t.Context(), t, client, tc.key, &tc.version) + var sdkErr *codersdk.Error + require.ErrorAs(t, err, &sdkErr) + require.Equal(t, tc.wantStatus, sdkErr.StatusCode()) + require.Contains(t, sdkErr.Error(), tc.wantMessage) + if tc.forbidErrMessage != "" { + require.NotContains(t, sdkErr.Error(), tc.forbidErrMessage) + } }) } } @@ -184,37 +212,90 @@ func TestAIGatewayServeMissingEntitlement(t *testing.T) { }) ctx := testutil.Context(t, testutil.WaitLong) - _, status := dialAIGatewayServe(ctx, t, client, "any-key", aibridgedproto.CurrentVersion.String()) - require.Equal(t, http.StatusForbidden, status) + // The production dialer must surface the upgrade failure as a + // *codersdk.Error so the standalone gateway's connect loop can detect the + // 403 and stop retrying instead of looping forever. + _, err := dialAIGatewayServe(ctx, t, client, "any-key") + var sdkErr *codersdk.Error + require.ErrorAs(t, err, &sdkErr) + require.Equal(t, http.StatusForbidden, sdkErr.StatusCode()) } -func TestAIGatewayServeDeletedKeyClosesActiveSession(t *testing.T) { +func TestAIGatewayServeTrackKeyUsageClosesActiveSession(t *testing.T) { t.Parallel() + t.Run("DeletedKey", func(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitShort) + session := setupActiveAIGatewayServeSession(ctx, t) + + //nolint:gocritic // Owner role is needed for gateway key management. + require.NoError(t, session.client.DeleteAIGatewayKey(ctx, session.created.ID)) + requireAIGatewayServeSessionClosed(t, session) + }) + + t.Run("RevokedEntitlement", func(t *testing.T) { + t.Parallel() + + ctx := testutil.Context(t, testutil.WaitShort) + session := setupActiveAIGatewayServeSession(ctx, t) + + licenses, err := session.client.Licenses(ctx) + require.NoError(t, err) + for _, license := range licenses { + require.NoError(t, session.client.DeleteLicense(ctx, license.ID)) + } + require.Eventually(t, func() bool { + return !session.api.Entitlements.Enabled(codersdk.FeatureAIBridge) + }, testutil.WaitShort, testutil.IntervalFast) + + requireAIGatewayServeSessionClosed(t, session) + }) +} + +type activeAIGatewayServeSession struct { + client *codersdk.Client + api *entcoderd.API + created codersdk.CreateAIGatewayKeyResponse + tick chan time.Time + dc aibridged.DRPCClient +} + +func setupActiveAIGatewayServeSession(ctx context.Context, t *testing.T) activeAIGatewayServeSession { + t.Helper() + tick := make(chan time.Time, 1) opts := aibridgeOpts(t) opts.Options.NewTicker = func(time.Duration) (<-chan time.Time, func()) { return tick, func() {} } - client, _ := coderdenttest.New(t, opts) - ctx := testutil.Context(t, testutil.WaitLong) + client, _, api, _ := coderdenttest.NewWithAPI(t, opts) //nolint:gocritic // Owner role is needed for gateway key management. - created, err := client.CreateAIGatewayKey(ctx, codersdk.CreateAIGatewayKeyRequest{Name: "serve-delete-active"}) + created, err := client.CreateAIGatewayKey(ctx, codersdk.CreateAIGatewayKeyRequest{Name: "key-name"}) require.NoError(t, err) - session, status := dialAIGatewayServe(ctx, t, client, created.Key, aibridgedproto.CurrentVersion.String()) - require.Equal(t, http.StatusSwitchingProtocols, status) - require.NotNil(t, session) + dc, err := dialAIGatewayServe(ctx, t, client, created.Key) + require.NoError(t, err) - //nolint:gocritic // Owner role is needed for gateway key management. - require.NoError(t, client.DeleteAIGatewayKey(ctx, created.ID)) + return activeAIGatewayServeSession{ + client: client, + api: api, + created: created, + tick: tick, + dc: dc, + } +} + +func requireAIGatewayServeSessionClosed(t *testing.T, s activeAIGatewayServeSession) { + t.Helper() - tick <- time.Now() // trigger aiGatewayTrackKeyUsage. + s.tick <- time.Now() // trigger gateway key / license check. require.Eventually(t, func() bool { select { - case <-session.CloseChan(): + case <-s.dc.DRPCConn().Closed(): return true default: return false