feat: add ai-gateway start command (#26605)

> AI Tools were used to produce this PR

This PR adds `coder ai-gateway start` command that runs the AI Gateway
as an independent process.

- Standalone process doesn't have access to DB. Uses DRPC services under
`/api/v2/ai-gateway/serve`for auth, recording and provider
initialization.
- It only handles LLM traffic, other endpoints (eg. `/sessions`) are
only available though `coderd`.
- The standalone gateway reuses applicable flags from AI Gateway
deployment options. Provider-seeding and coderd-only options are
excluded.
- Only added to fat build, the slim build stub rejects the command.

Some wiring used by this new command is added.

**`NewWebsocketDialer`** - implements the standalone gateway's
connection to coderd's `/api/v2/ai-gateway/serve` endpoint. It upgrades
to a WebSocket, multiplexes with yamux, and wires all DRPC services.

**`AIGatewayDataPlaneMiddleware`** - extracts the per-request middleware
chain (concurrency limiting, rate limiting, BYOK gating) into a shared
function used by both the embedded route and the standalone gateway.

**`RootCmd.ResolveClientConnection`** - resolve the deployment URL and
builds an HTTP transport without requiring a session token. Used in
`ai-gateway start`command as it authenticates using different credential
type.

---------

Co-authored-by: Danny Kopping <danny@coder.com>
This commit is contained in:
Paweł Banaszewski
2026-07-08 11:12:53 +02:00
committed by GitHub
co-authored by Danny Kopping
parent 195dffc651
commit ccba3969ab
18 changed files with 1418 additions and 150 deletions
+1
View File
@@ -20,6 +20,7 @@ func (r *RootCmd) aiGateway() *serpent.Command {
return inv.Command.HelpHandler(inv)
},
Children: []*serpent.Command{
r.aiGatewayStart(),
r.aiGatewayKeys(),
},
}
+309
View File
@@ -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()
}
@@ -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)
}
+24
View File
@@ -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
}
+2 -1
View File
@@ -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.
@@ -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.
+30 -19
View File
@@ -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
+15 -5
View File
@@ -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
+173 -92
View File
@@ -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) {
type versionOverridingRoundTripper struct {
baseTransport http.RoundTripper
overrideAPIVersion string
}
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)
}
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)
}
func dialAIGatewayServeWithVersion(ctx context.Context, t *testing.T, client *codersdk.Client, key string, version *string) (aibridged.DRPCClient, error) {
t.Helper()
serverURL, err := client.URL.Parse("/api/v2/ai-gateway/serve")
require.NoError(t, err)
query := serverURL.Query()
if version != "" {
query.Set("version", version)
}
serverURL.RawQuery = query.Encode()
headers := http.Header{}
if key != "" {
headers.Set(codersdk.AIGatewayKeyHeader, key)
}
conn, res, err := websocket.Dial(ctx, serverURL.String(), &websocket.DialOptions{
HTTPClient: &http.Client{Transport: client.HTTPClient.Transport},
CompressionMode: websocket.CompressionDisabled,
HTTPHeader: headers,
})
if err != nil {
statusCode := 0
if res != nil {
statusCode = res.StatusCode
_ = res.Body.Close()
transport := client.HTTPClient.Transport
if version != nil {
transport = versionOverridingRoundTripper{
baseTransport: transport,
overrideAPIVersion: *version,
}
return nil, statusCode
}
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)
dc, err := aibridged.NewWebsocketDialer(client.URL, transport, key)(ctx)
if err != nil {
return nil, 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,
}
}
tick <- time.Now() // trigger aiGatewayTrackKeyUsage.
func requireAIGatewayServeSessionClosed(t *testing.T, s activeAIGatewayServeSession) {
t.Helper()
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