mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
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:
@@ -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(¶ms)
|
||||
}
|
||||
|
||||
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()
|
||||
|
||||
@@ -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"))
|
||||
}
|
||||
Reference in New Issue
Block a user