test(enterprise/cli): add standalone AI Gateway connection tests (#27860)

Adds two tests that run the real `ai-gateway start` command against a
real coderd over the production websocket dialer.
`TestAIGatewayStartE2E`: the gateway completes the handshake, loads
providers over DRPC, proxies an OpenAI chat completion on its own
listener, and the interception is recorded in coderd.
`TestAIGatewayStartE2E_InvalidKey`: a key rejected by the handshake is
fatal rather than retried, and the command reports it.

Also tidies the existing tests: `TestAIGatewayStart_HealthBeforeReady`
moves to the external package and reuses the new helpers, the two fake
reloaders collapse into one `mockReloader`.

---

Generated with Coder Agents.
This commit is contained in:
Paweł Banaszewski
2026-08-17 20:48:09 +00:00
committed by GitHub
parent 46ec620767
commit 94f487b890
2 changed files with 361 additions and 225 deletions
+120 -225
View File
@@ -5,7 +5,6 @@ package cli
import (
"context"
"errors"
"fmt"
"net"
"net/http"
"net/http/httptest"
@@ -21,42 +20,32 @@ import (
"storj.io/drpc"
"cdr.dev/slog/v3"
"github.com/coder/coder/v2/cli/clitest"
agplaibridge "github.com/coder/coder/v2/coderd/aibridge"
"github.com/coder/coder/v2/coderd/aibridged"
"github.com/coder/coder/v2/codersdk"
"github.com/coder/coder/v2/testutil"
)
// 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 {
// mockReloader returns err for the first failUntil Reload calls, then nil.
// onReload runs at the end of every reload, letting a test hang the retry loop.
type mockReloader struct {
calls atomic.Int32
failUntil int32
err error
onReload func()
}
func (r *failThenSucceedReloader) Reload(_ context.Context) error {
if r.calls.Add(1) <= r.failUntil {
return xerrors.New("transient failure")
func (r *mockReloader) Reload(context.Context) error {
failed := r.calls.Add(1) <= r.failUntil
if r.onReload != nil {
r.onReload()
}
if failed {
return r.err
}
return nil
}
type failingReloader struct {
after func()
calls atomic.Int32
err error
}
func (r *failingReloader) Reload(context.Context) error {
r.calls.Add(1)
if r.after != nil {
r.after()
}
return r.err
}
type connectedDRPCConn struct {
drpc.Conn
closed chan struct{}
@@ -95,12 +84,14 @@ func (p *controlledShutdownPool) Shutdown(ctx context.Context) error {
return errors.Join(p.CachedBridgePool.Shutdown(ctx), p.err)
}
type standaloneGatewayTestParams struct {
params standaloneGatewayParams
pool *controlledShutdownPool
}
// testGatewayOption mutates the default standaloneGatewayParams before the
// gateway is constructed.
type testGatewayOption func(*standaloneGatewayParams)
func newStandaloneGatewayTestParams(t *testing.T) *standaloneGatewayTestParams {
// newTestStandaloneGateway constructs a standalone gateway for
// testing, with a controllable shutdown pool and optional customizations.
// Uses blockingStandaloneDaemonDialer to mock coderd.
func newTestStandaloneGateway(t *testing.T, opts ...testGatewayOption) (*standaloneGateway, *controlledShutdownPool) {
t.Helper()
logger := slog.Make()
@@ -112,18 +103,26 @@ func newStandaloneGatewayTestParams(t *testing.T) *standaloneGatewayTestParams {
})
pool := &controlledShutdownPool{CachedBridgePool: cachedPool}
return &standaloneGatewayTestParams{
params: standaloneGatewayParams{
httpAddress: "127.0.0.1:0",
params := standaloneGatewayParams{
httpAddress: "127.0.0.1:0",
dialer: blockingStandaloneDaemonDialer,
pool: pool,
dialer: blockingStandaloneDaemonDialer,
pool: pool,
logger: logger,
tracer: tracer,
},
pool: pool,
logger: logger,
tracer: tracer,
}
for _, m := range opts {
m(&params)
}
gateway, err := newStandaloneGateway(params)
require.NoError(t, err)
t.Cleanup(func() {
require.NoError(t, shutdownWithTimeout(gateway.daemon.Shutdown, testutil.WaitShort))
})
return gateway, pool
}
func TestStandaloneGatewayLoadProviders(t *testing.T) {
@@ -131,40 +130,37 @@ func TestStandaloneGatewayLoadProviders(t *testing.T) {
reloadErr := xerrors.New("reload failed")
tests := []struct {
name string
setup func(*testing.T, *aibridged.Server, context.CancelFunc) (aibridged.ProviderReloader, *atomic.Int32)
wantErr error
wantCalls int32
wantLoaded bool
name string
reloaderFailUntil int32
reloaderErr error
reloaderAfter func(t *testing.T, daemon *aibridged.Server, cancel context.CancelFunc)
wantErr error
wantCalls int32
wantLoaded bool
}{
{
name: "Retry succeeds",
setup: func(_ *testing.T, _ *aibridged.Server, _ context.CancelFunc) (aibridged.ProviderReloader, *atomic.Int32) {
reloader := &failThenSucceedReloader{failUntil: 2}
return reloader, &reloader.calls
},
wantCalls: 3,
wantLoaded: true,
name: "Retry succeeds",
reloaderFailUntil: 2,
reloaderErr: xerrors.New("transient failure"),
wantCalls: 3,
wantLoaded: true,
},
{
name: "Daemon stops retry",
setup: func(t *testing.T, daemon *aibridged.Server, _ context.CancelFunc) (aibridged.ProviderReloader, *atomic.Int32) {
reloader := &failingReloader{
after: func() {
require.NoError(t, daemon.Close())
},
err: reloadErr,
}
return reloader, &reloader.calls
name: "Daemon stops retry",
reloaderFailUntil: 1,
reloaderErr: reloadErr,
reloaderAfter: func(t *testing.T, daemon *aibridged.Server, _ context.CancelFunc) {
require.NoError(t, daemon.Close())
},
wantErr: reloadErr,
wantCalls: 1,
},
{
name: "Context cancellation stops retry",
setup: func(_ *testing.T, _ *aibridged.Server, cancel context.CancelFunc) (aibridged.ProviderReloader, *atomic.Int32) {
reloader := &failingReloader{after: cancel, err: reloadErr}
return reloader, &reloader.calls
name: "Context cancellation stops retry",
reloaderFailUntil: 1,
reloaderErr: reloadErr,
reloaderAfter: func(_ *testing.T, _ *aibridged.Server, cancel context.CancelFunc) {
cancel()
},
wantErr: context.Canceled,
wantCalls: 1,
@@ -177,13 +173,13 @@ func TestStandaloneGatewayLoadProviders(t *testing.T) {
ctx, cancel := context.WithCancel(testutil.Context(t, testutil.WaitShort))
defer cancel()
logger := slog.Make()
daemon := newTestStandaloneDaemon(t, logger)
reloader, calls := tc.setup(t, daemon, cancel)
gateway := &standaloneGateway{
daemon: daemon,
providerLogger: logger,
reloader: reloader,
reloader := &mockReloader{failUntil: tc.reloaderFailUntil, err: tc.reloaderErr}
gateway, _ := newTestStandaloneGateway(t)
gateway.reloader = reloader
// reloaderAfter needs daemon which only exists once the gateway is constructed,
// nothing reloads until loadProviders below.
if tc.reloaderAfter != nil {
reloader.onReload = func() { tc.reloaderAfter(t, gateway.daemon, cancel) }
}
err := gateway.loadProviders(ctx)
@@ -192,7 +188,7 @@ func TestStandaloneGatewayLoadProviders(t *testing.T) {
} else {
require.ErrorIs(t, err, tc.wantErr)
}
require.Equal(t, tc.wantCalls, calls.Load())
require.Equal(t, tc.wantCalls, reloader.calls.Load())
require.Equal(t, tc.wantLoaded, gateway.providersLoaded.Load())
})
}
@@ -202,34 +198,19 @@ func TestStandaloneGatewayHealthAndReadiness(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitShort)
logger := slog.Make()
tracer := sdktrace.NewTracerProvider().Tracer("test")
pool, err := aibridged.NewCachedBridgePool(aibridged.DefaultPoolOptions, nil, logger, nil, tracer)
require.NoError(t, err)
connections := make(chan drpc.Conn, 2)
dialer := func(ctx context.Context) (aibridged.DRPCClient, error) {
select {
case conn := <-connections:
return &aibridged.Client{Conn: conn}, nil
case <-ctx.Done():
return nil, ctx.Err()
modifyDialer := func(p *standaloneGatewayParams) {
p.dialer = func(ctx context.Context) (aibridged.DRPCClient, error) {
select {
case conn := <-connections:
return &aibridged.Client{Conn: conn}, nil
case <-ctx.Done():
return nil, ctx.Err()
}
}
}
daemon, err := aibridged.New(ctx, pool, dialer, logger, tracer)
require.NoError(t, err)
t.Cleanup(func() {
require.NoError(t, shutdownWithTimeout(daemon.Shutdown, testutil.WaitShort))
})
gateway := &standaloneGateway{
daemon: daemon,
providerLogger: logger,
reloader: &failThenSucceedReloader{},
}
gateway.httpServer = &http.Server{
Handler: newGatewayMux(daemon, gateway.ready, func(next http.Handler) http.Handler { return next }),
ReadHeaderTimeout: testutil.WaitShort,
}
gateway, _ := newTestStandaloneGateway(t, modifyDialer)
gateway.reloader = &mockReloader{}
// The HTTP server is healthy before the daemon connects or providers load.
require.Equal(t, http.StatusOK, healthzStatus(t, gateway))
@@ -238,7 +219,7 @@ func TestStandaloneGatewayHealthAndReadiness(t *testing.T) {
// A daemon connection alone does not make the gateway ready.
firstConn := &connectedDRPCConn{closed: make(chan struct{})}
connections <- firstConn
require.Eventually(t, daemon.Ready, testutil.WaitShort, testutil.IntervalFast)
require.Eventually(t, gateway.daemon.Ready, testutil.WaitShort, testutil.IntervalFast)
require.Equal(t, http.StatusOK, healthzStatus(t, gateway))
require.Equal(t, http.StatusServiceUnavailable, readyzStatus(t, gateway))
@@ -249,13 +230,13 @@ func TestStandaloneGatewayHealthAndReadiness(t *testing.T) {
// Losing the daemon connection affects readiness but not HTTP health.
require.NoError(t, firstConn.Close())
require.Eventually(t, func() bool { return !daemon.Ready() }, testutil.WaitShort, testutil.IntervalFast)
require.Eventually(t, func() bool { return !gateway.daemon.Ready() }, testutil.WaitShort, testutil.IntervalFast)
require.Equal(t, http.StatusOK, healthzStatus(t, gateway))
require.Equal(t, http.StatusServiceUnavailable, readyzStatus(t, gateway))
// Readiness recovers when the daemon reconnects; providers remain loaded.
connections <- &connectedDRPCConn{closed: make(chan struct{})}
require.Eventually(t, daemon.Ready, testutil.WaitShort, testutil.IntervalFast)
require.Eventually(t, gateway.daemon.Ready, testutil.WaitShort, testutil.IntervalFast)
require.Equal(t, http.StatusOK, healthzStatus(t, gateway))
require.Equal(t, http.StatusOK, readyzStatus(t, gateway))
}
@@ -277,73 +258,13 @@ func probeStatus(t *testing.T, gateway *standaloneGateway, path string) int {
return rec.Code
}
func TestAIGatewayStart_HealthBeforeReady(t *testing.T) {
t.Parallel()
// Fake coderd that answers 503 so the daemon keeps retrying to connect.
coderSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusServiceUnavailable)
}))
t.Cleanup(coderSrv.Close)
gatewayAddress := fmt.Sprintf("127.0.0.1:%d", testutil.RandomPort(t))
var root RootCmd
cmd, err := root.Command(root.enterpriseOnly())
require.NoError(t, err)
inv, _ := clitest.NewWithCommand(t, cmd,
"--url", coderSrv.URL,
"ai-gateway", "start",
"--key", "test-key",
"--http-address", gatewayAddress,
)
ctx := testutil.Context(t, testutil.WaitShort)
// Watch the command for an early exit so a clash on gatewayAddress is
// reported as a bind failure instead of a probe timeout.
cmdDone := make(chan error, 1)
clitest.StartWithAssert(t, inv.WithContext(ctx), func(_ *testing.T, err error) {
cmdDone <- err
})
// healthz check
client := &http.Client{Timeout: testutil.WaitShort}
baseURL := "http://" + gatewayAddress
testutil.Eventually(ctx, t, func(ctx context.Context) bool {
select {
case err := <-cmdDone:
t.Fatalf("ai-gateway start exited before serving: %v", err)
default:
}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, baseURL+healthzPath, nil)
if err != nil {
return false
}
resp, err := client.Do(req)
if err != nil {
return false
}
defer resp.Body.Close()
return resp.StatusCode == http.StatusOK
}, testutil.IntervalFast)
// readyz check (unavailable due to no connection to coderd)
req, err := http.NewRequestWithContext(ctx, http.MethodGet, baseURL+readyzPath, nil)
require.NoError(t, err)
resp, err := client.Do(req)
require.NoError(t, err)
defer resp.Body.Close()
require.Equal(t, http.StatusServiceUnavailable, resp.StatusCode)
}
func TestRunStandaloneGateway_ContextCanceled(t *testing.T) {
t.Parallel()
testCtx := testutil.Context(t, testutil.WaitShort)
runCtx, cancelRun := context.WithCancel(testCtx)
defer cancelRun()
test := newStandaloneGatewayTestParams(t)
gateway, err := newStandaloneGateway(test.params)
require.NoError(t, err)
gateway, _ := newTestStandaloneGateway(t)
runDone := make(chan error, 1)
go func() {
runDone <- gateway.run(runCtx)
@@ -360,14 +281,13 @@ func TestRunStandaloneGateway_ContextCanceled(t *testing.T) {
func TestRunStandaloneGateway_DaemonExited(t *testing.T) {
t.Parallel()
test := newStandaloneGatewayTestParams(t)
test.params.dialer = func(context.Context) (aibridged.DRPCClient, error) {
return nil, codersdk.NewError(http.StatusUnauthorized, codersdk.Response{Message: "invalid gateway key"})
modifyDialer := func(p *standaloneGatewayParams) {
p.dialer = func(context.Context) (aibridged.DRPCClient, error) {
return nil, codersdk.NewError(http.StatusUnauthorized, codersdk.Response{Message: "invalid gateway key"})
}
}
gateway, err := newStandaloneGateway(test.params)
require.NoError(t, err)
err = gateway.run(testutil.Context(t, testutil.WaitShort))
gateway, _ := newTestStandaloneGateway(t, modifyDialer)
err := gateway.run(testutil.Context(t, testutil.WaitShort))
require.ErrorContains(t, err, "AI Gateway daemon exited")
require.True(t, gateway.drpcClosed.Load(), "DRPC connection must be closed before run returns")
require.True(t, gateway.providerRefreshStopped.Load(), "provider refresh must stop before run returns")
@@ -378,18 +298,18 @@ func TestRunStandaloneGateway_HTTPStopsBeforeDaemonShutdown(t *testing.T) {
t.Parallel()
testCtx := testutil.Context(t, testutil.WaitShort)
test := newStandaloneGatewayTestParams(t)
modifyTLS := func(p *standaloneGatewayParams) {
p.tlsCertFile = filepath.Join(t.TempDir(), "missing.crt")
p.tlsKeyFile = filepath.Join(t.TempDir(), "missing.key")
}
gateway, pool := newTestStandaloneGateway(t, modifyTLS)
shutdownErr := xerrors.New("pool shutdown failed")
shutdownStarted := make(chan struct{}, 1)
shutdownRelease := make(chan struct{})
test.pool.err = shutdownErr
test.pool.started = shutdownStarted
test.pool.release = shutdownRelease
test.params.tlsCertFile = filepath.Join(t.TempDir(), "missing.crt")
test.params.tlsKeyFile = filepath.Join(t.TempDir(), "missing.key")
pool.err = shutdownErr
pool.started = shutdownStarted
pool.release = shutdownRelease
gateway, err := newStandaloneGateway(test.params)
require.NoError(t, err)
runDone := make(chan error, 1)
go func() {
runDone <- gateway.run(testCtx)
@@ -400,7 +320,7 @@ func TestRunStandaloneGateway_HTTPStopsBeforeDaemonShutdown(t *testing.T) {
require.False(t, gateway.drpcClosed.Load(), "DRPC connection must remain open until daemon shutdown completes")
close(shutdownRelease)
err = testutil.RequireReceive(testCtx, t, runDone)
err := testutil.RequireReceive(testCtx, t, runDone)
require.False(t, gateway.drpcClosed.Load(), "DRPC connection shutdown must not be marked successful after an error")
require.ErrorContains(t, err, "serve:")
require.ErrorContains(t, err, "shutdown AI Gateway daemon:")
@@ -410,17 +330,20 @@ func TestRunStandaloneGateway_HTTPStopsBeforeDaemonShutdown(t *testing.T) {
func TestRunStandaloneGateway_ListenAndShutdownErrors(t *testing.T) {
t.Parallel()
test := newStandaloneGatewayTestParams(t)
shutdownErr := xerrors.New("pool shutdown failed")
test.pool.err = shutdownErr
listener, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
t.Cleanup(func() {
_ = listener.Close()
})
test.params.httpAddress = listener.Addr().String()
// Occupy an address so binding the gateway listener fails.
modifyAddr := func(p *standaloneGatewayParams) {
p.httpAddress = listener.Addr().String()
}
gateway, pool := newTestStandaloneGateway(t, modifyAddr)
shutdownErr := xerrors.New("pool shutdown failed")
pool.err = shutdownErr
err = runStandaloneGateway(testutil.Context(t, testutil.WaitShort), test.params)
err = gateway.run(testutil.Context(t, testutil.WaitShort))
require.NoError(t, listener.Close())
require.ErrorContains(t, err, "listen on")
require.ErrorContains(t, err, "shutdown AI Gateway daemon:")
@@ -432,45 +355,31 @@ func TestStandaloneGatewayServe_ShutdownOrder(t *testing.T) {
// Set up a running daemon, provider reloader, and blocked HTTP request.
testCtx := testutil.Context(t, testutil.WaitShort)
logger := slog.Make()
pool, err := aibridged.NewCachedBridgePool(aibridged.DefaultPoolOptions, nil, logger, nil, sdktrace.NewTracerProvider().Tracer("test"))
require.NoError(t, err)
daemon, err := aibridged.New(context.Background(), pool, blockingStandaloneDaemonDialer, logger, sdktrace.NewTracerProvider().Tracer("test"))
require.NoError(t, err)
httpAddress := "127.0.0.1:0"
// inFlightPath scopes the blocking handler to the request this test keeps in
// flight. Another test may probe a port the OS later assigns to this
// listener, and such traffic must not stand in for that request.
const inFlightPath = "/in-flight"
reloader := &failThenSucceedReloader{}
reloader := &mockReloader{}
handlerStarted := make(chan struct{}, 1)
httpShutdownStarted := make(chan struct{}, 1)
releaseHandler := make(chan struct{})
gateway := &standaloneGateway{
daemon: daemon,
httpServer: &http.Server{
Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != inFlightPath {
w.WriteHeader(http.StatusNotFound)
return
}
select {
case handlerStarted <- struct{}{}:
default:
}
<-releaseHandler
w.WriteHeader(http.StatusNoContent)
}),
ReadHeaderTimeout: testutil.WaitShort,
},
httpAddress: httpAddress,
logger: logger,
providerLogger: logger,
reloader: reloader,
listenerReady: make(chan struct{}),
}
gateway, _ := newTestStandaloneGateway(t)
gateway.reloader = reloader
// The gateway mux is replaced with a handler this test can block, so an
// in-flight request is observable during shutdown.
gateway.httpServer.Handler = http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != inFlightPath {
w.WriteHeader(http.StatusNotFound)
return
}
select {
case handlerStarted <- struct{}{}:
default:
}
<-releaseHandler
w.WriteHeader(http.StatusNoContent)
})
gateway.httpServer.RegisterOnShutdown(func() {
httpShutdownStarted <- struct{}{}
})
@@ -539,20 +448,6 @@ func requireListening(ctx context.Context, t *testing.T, gateway *standaloneGate
return ""
}
func newTestStandaloneDaemon(t *testing.T, logger slog.Logger) *aibridged.Server {
t.Helper()
tracer := sdktrace.NewTracerProvider().Tracer("test")
pool, err := aibridged.NewCachedBridgePool(aibridged.DefaultPoolOptions, nil, logger, nil, tracer)
require.NoError(t, err)
daemon, err := aibridged.New(context.Background(), pool, blockingStandaloneDaemonDialer, logger, tracer)
require.NoError(t, err)
t.Cleanup(func() {
require.NoError(t, daemon.Close())
})
return daemon
}
func blockingStandaloneDaemonDialer(ctx context.Context) (aibridged.DRPCClient, error) {
<-ctx.Done()
return nil, ctx.Err()
+241
View File
@@ -0,0 +1,241 @@
//go:build !slim
package cli_test
import (
"context"
"io"
"net/http"
"net/http/httptest"
"regexp"
"strings"
"sync/atomic"
"testing"
"github.com/stretchr/testify/require"
"github.com/coder/coder/v2/cli/clitest"
"github.com/coder/coder/v2/coderd/coderdtest"
"github.com/coder/coder/v2/codersdk"
"github.com/coder/coder/v2/enterprise/coderd/coderdenttest"
"github.com/coder/coder/v2/enterprise/coderd/license"
"github.com/coder/coder/v2/pty/ptytest"
"github.com/coder/coder/v2/testutil"
"github.com/coder/serpent"
)
// The end-to-end tests in this file run the real `ai-gateway start` command
// against a real coderd and only observe the public surface: CLI flags, log
// output, the gateway's HTTP endpoints, coderd's API, and the command's exit
// error. Each part of the standalone plumbing executes at least once. Detailed
// behavior (readiness transitions, reconnect semantics, shutdown ordering) is
// covered by the internal tests in aigatewaystart_internal_test.go and the
// reconnection tests, which construct the gateway directly.
// aiGatewayChatCompletionRequest is the LLM request sent through the gateway. The
// model is asserted against the recorded interception.
const aiGatewayChatCompletionRequest = `{"messages":[{"role":"user","content":"standalone gateway e2e"}],"model":"gpt-4.1"}`
// aiGatewayUpstreamResponse is the fixed completion the mock upstream LLM API
// returns, so the test can recognize it in the gateway's response.
const aiGatewayUpstreamResponse = `{
"id": "chatcmpl-e2e",
"object": "chat.completion",
"created": 1753343279,
"model": "gpt-4.1",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "standalone gateway e2e response"},
"finish_reason": "stop"
}
],
"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}
}`
// aiGatewayDeployment is a coderd entitled for the AI Gateway, with a gateway key, a
// configured provider backed by a mock upstream, and a member user whose
// session token authenticates LLM traffic.
type aiGatewayDeployment struct {
client *codersdk.Client
userClient *codersdk.Client
user codersdk.User
key string
upstreamHits *atomic.Int32
}
func setupAIGatewayCoderdenttestDeployment(t *testing.T) (*codersdk.Client, codersdk.CreateFirstUserResponse) {
dv := coderdtest.DeploymentValues(t)
dv.AI.BridgeConfig.Enabled = serpent.Bool(true)
return coderdenttest.New(t, &coderdenttest.Options{
Options: &coderdtest.Options{DeploymentValues: dv},
LicenseOptions: &coderdenttest.LicenseOptions{
Features: license.Features{codersdk.FeatureAIBridge: 1},
},
})
}
func setupAIGatewayDeployment(ctx context.Context, t *testing.T) *aiGatewayDeployment {
t.Helper()
client, firstUser := setupAIGatewayCoderdenttestDeployment(t)
//nolint:gocritic // Owner role is needed for gateway key management.
key, err := client.CreateAIGatewayKey(ctx, codersdk.CreateAIGatewayKeyRequest{Name: "e2e"})
require.NoError(t, err)
var hits atomic.Int32
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
hits.Add(1)
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusOK)
_, _ = w.Write([]byte(aiGatewayUpstreamResponse))
}))
t.Cleanup(upstream.Close)
//nolint:gocritic // Owner role is needed for provider management.
_, err = client.CreateAIProvider(ctx, codersdk.CreateAIProviderRequest{
Type: codersdk.AIProviderTypeOpenAI,
Name: "openai",
Enabled: true,
BaseURL: upstream.URL,
APIKeys: []string{"sk-e2e"},
})
require.NoError(t, err)
userClient, user := coderdtest.CreateAnotherUser(t, client, firstUser.OrganizationID)
return &aiGatewayDeployment{
client: client,
userClient: userClient,
user: user,
key: key.Key,
upstreamHits: &hits,
}
}
// startAIGatewayCommand runs `ai-gateway start` and returns the base URL of
// its HTTP listener, discovered from the startup log line, together with the
// command's error waiter.
func startAIGatewayCommand(ctx context.Context, t *testing.T, coderURL, key string) (string, *clitest.ErrorWaiter) {
t.Helper()
inv, _ := newCLI(t,
"ai-gateway", "start",
"--url", coderURL,
"--key", key,
"--http-address", "127.0.0.1:0",
)
inv = inv.WithContext(ctx)
pty := ptytest.New(t).Attach(inv)
waiter := clitest.StartWithWaiter(t, inv)
// Extract bound address from the startup log.
pty.ExpectMatch(ctx, "standalone AI Gateway listening")
line := pty.ReadLine(ctx)
matches := regexp.MustCompile(`address=([0-9.]+:[0-9]+)`).FindStringSubmatch(line)
require.Len(t, matches, 2, "listener address not found in startup log: %q", line)
return "http://" + matches[1], waiter
}
func getAIGatewayStatus(ctx context.Context, t *testing.T, url string) int {
t.Helper()
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
require.NoError(t, err)
resp, err := http.DefaultClient.Do(req)
require.NoError(t, err)
defer resp.Body.Close()
return resp.StatusCode
}
// TestAIGatewayStartE2E drives every part of the standalone gateway plumbing
// once through public surface only: the CLI starts with flags, connects to
// coderd with a gateway key, reports health and readiness, proxies an LLM
// request from a real client to a real upstream, records the interception in
// coderd, and shuts down cleanly.
func TestAIGatewayStartE2E(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
dep := setupAIGatewayDeployment(ctx, t)
baseURL, waiter := startAIGatewayCommand(ctx, t, dep.client.URL.String(), dep.key)
// Liveness holds as soon as the listener is up; readiness follows once the
// DRPC connection is established and providers are loaded.
require.Equal(t, http.StatusOK, getAIGatewayStatus(ctx, t, baseURL+"/healthz"))
require.Eventually(t, func() bool {
return getAIGatewayStatus(ctx, t, baseURL+"/readyz") == http.StatusOK
}, testutil.WaitLong, testutil.IntervalFast)
// One LLM request through the gateway's own listener.
req, err := http.NewRequestWithContext(ctx, http.MethodPost,
baseURL+"/openai/v1/chat/completions", strings.NewReader(aiGatewayChatCompletionRequest))
require.NoError(t, err)
req.Header.Set("Authorization", "Bearer "+dep.userClient.SessionToken())
req.Header.Set("Content-Type", "application/json")
resp, err := http.DefaultClient.Do(req)
require.NoError(t, err)
body, err := io.ReadAll(resp.Body)
_ = resp.Body.Close()
require.NoError(t, err)
require.Equal(t, http.StatusOK, resp.StatusCode, "body: %s", body)
require.Contains(t, string(body), "standalone gateway e2e response")
require.Equal(t, int32(1), dep.upstreamHits.Load())
// The interception is recorded in coderd. Recording is asynchronous, so
// the assertion has to be eventual.
require.Eventually(t, func() bool {
//nolint:gocritic // Owner role is needed to list every user's sessions.
sessions, err := dep.client.AIBridgeListSessions(ctx, codersdk.AIBridgeListSessionsFilter{
Initiator: dep.user.Username,
})
return err == nil && len(sessions.Sessions) == 1
}, testutil.WaitLong, testutil.IntervalFast)
// Graceful shutdown: canceling the command must produce a clean exit.
waiter.Cancel()
require.NoError(t, waiter.Wait())
}
// TestAIGatewayStartE2E_InvalidKey covers the fatal error plumbing: a gateway
// started with a key coderd rejects must exit with the rejection instead of
// retrying forever.
func TestAIGatewayStartE2E_InvalidKey(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitLong)
client, _ := setupAIGatewayCoderdenttestDeployment(t)
inv, _ := newCLI(t,
"ai-gateway", "start",
"--url", client.URL.String(),
"--key", "not-a-valid-key",
"--http-address", "127.0.0.1:0",
)
inv = inv.WithContext(ctx)
waiter := clitest.StartWithWaiter(t, inv)
waiter.RequireContains("AI Gateway key invalid")
}
// TestAIGatewayStart_HealthBeforeReady covers the split between liveness and
// readiness: the listener serves /healthz as soon as it is bound, while
// /readyz stays 503 until the daemon reaches coderd.
func TestAIGatewayStart_HealthBeforeReady(t *testing.T) {
t.Parallel()
// Fake coderd that answers 503 so the daemon keeps retrying to connect.
coderSrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusServiceUnavailable)
}))
t.Cleanup(coderSrv.Close)
ctx := testutil.Context(t, testutil.WaitShort)
baseURL, _ := startAIGatewayCommand(ctx, t, coderSrv.URL, "test-key")
// The startup log line is emitted after the listener is bound, so no retry
// loop is needed.
require.Equal(t, http.StatusOK, getAIGatewayStatus(ctx, t, baseURL+"/healthz"))
require.Equal(t, http.StatusServiceUnavailable, getAIGatewayStatus(ctx, t, baseURL+"/readyz"))
}