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