mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
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:
co-authored by
Danny Kopping
parent
195dffc651
commit
ccba3969ab
@@ -20,6 +20,7 @@ func (r *RootCmd) aiGateway() *serpent.Command {
|
||||
return inv.Command.HelpHandler(inv)
|
||||
},
|
||||
Children: []*serpent.Command{
|
||||
r.aiGatewayStart(),
|
||||
r.aiGatewayKeys(),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
@@ -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.
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user