fix: remove standalone AI Gateway http listener dependency on loading providers (#27303)

Fixes an issue where the standalone AI Gateway waited for the initial
provider load before starting its HTTP server.

HTTP serving now starts independently of provider synchronization.
`/healthz` becomes available when the HTTP server starts, while
`/readyz` requires an active DRPC connection and completed initial
provider load.

Enables the Helm chart's startup and liveness probes by default because
liveness no longer depends on provider loading.
This commit is contained in:
Paweł Banaszewski
2026-07-23 11:55:42 +02:00
committed by GitHub
parent b9fad66214
commit 468b1a27a3
8 changed files with 765 additions and 232 deletions
+2 -2
View File
@@ -90,7 +90,7 @@ func newAIBridgeDaemon(coderAPI *coderd.API, cfg codersdk.AIBridgeConfig, reg pr
// and the standalone gateway (WebSocket RPC, retried at startup) so the fetch,
// build, replace, and reload-metric accounting live in one place.
type poolRPCReloader struct {
pool *aibridged.CachedBridgePool
pool aibridged.Pooler
client aibridged.ClientFuncWithContext
cfg codersdk.AIBridgeConfig
logger slog.Logger
@@ -104,7 +104,7 @@ type poolRPCReloader struct {
// Reload's context, so a blocking acquisition unblocks when that context is
// canceled.
func NewPoolRPCReloader(
pool *aibridged.CachedBridgePool,
pool aibridged.Pooler,
client aibridged.ClientFuncWithContext,
cfg codersdk.AIBridgeConfig,
logger slog.Logger,
+246 -148
View File
@@ -11,6 +11,7 @@ import (
"os"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/prometheus/client_golang/prometheus"
@@ -35,11 +36,14 @@ import (
)
const (
// helm/ai-gateway's terminationGracePeriodSeconds must exceed
// shutdownTimeout so graceful shutdown completes before Kubernetes sends
// SIGKILL.
shutdownTimeout = 5 * time.Minute
traceShutdownTimeout = 5 * time.Second
// The sum of daemonShutdownTimeout, httpShutdownTimeout,
// providerReloadShutdownTimeout, and traceShutdownTimeout must stay below
// terminationGracePeriodSeconds in helm/ai-gateway/values.yaml so the
// process can complete graceful shutdown before Kubernetes sends SIGKILL.
daemonShutdownTimeout = 5 * time.Second
httpShutdownTimeout = 5 * time.Minute
providerReloadShutdownTimeout = 5 * time.Second
traceShutdownTimeout = 5 * time.Second
healthzPath = "/healthz"
readyzPath = "/readyz"
@@ -143,6 +147,9 @@ func (r *RootCmd) aiGatewayStart() *serpent.Command {
providerMetrics := aibridged.NewMetrics(registry)
tracerProvider, _, closeTracing := agpl.ConfigureTraceProviderWithService(signalCtx, logger, vals, "coder-ai-gateway")
// The tracer is shared by the gateway's HTTP middleware, pool, and
// daemon, so it must be flushed only after runStandaloneGateway returns
// and all span producers have stopped, hence the handler-level defer.
defer func() {
logger.Debug(signalCtx, "closing tracing")
traceCloseErr := shutdownWithTimeout(closeTracing, traceShutdownTimeout)
@@ -167,98 +174,21 @@ func (r *RootCmd) aiGatewayStart() *serpent.Command {
}
registry.MustRegister(keypool.NewStateCollector(pool.KeyPools))
dialer := aibridged.NewWebsocketDialer(serverURL, transport, resolvedKey)
aibridgedCtx, aibridgedCancel := context.WithCancel(context.Background())
defer aibridgedCancel()
srv, err := aibridged.New(aibridgedCtx, pool, dialer, gatewayLogger, tracer)
if err != nil {
return xerrors.Errorf("start AI Gateway daemon: %w", err)
}
defer srv.Close()
return runStandaloneGateway(signalCtx, standaloneGatewayParams{
bridgeConfig: vals.AI.BridgeConfig,
coderURL: serverURL.String(),
httpAddress: httpAddress,
tlsCertFile: tlsCertFile,
tlsKeyFile: tlsKeyFile,
// Fetch the initial provider set from coderd, retrying until
// success. Subsequent changes are delivered by the watch loop
// started below. The reloader's client acquisition honors the
// context of each Reload call, so loadProviders is bounded by
// signalCtx and the watch loop by watchCtx.
providerLogger := gatewayLogger.Named("providers")
reloader := agpl.NewPoolRPCReloader(pool, srv.ClientContext, 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)
}
dialer: aibridged.NewWebsocketDialer(serverURL, transport, resolvedKey),
pool: pool,
mw := gatewayMiddleware(vals.AI.BridgeConfig, tracer)
// Watch coderd for provider changes and refresh the pool on each
// signal.
watchCtx, watchCancel := context.WithCancel(signalCtx)
var watchWG sync.WaitGroup
watchWG.Go(func() {
// srv.ClientContext observes watchCtx, so watchCancel below
// unblocks a pending client acquisition and drains this
// goroutine without relying on srv.Close.
if err := aibridged.WatchProviderReload(watchCtx, srv.ClientContext, reloader, providerLogger); err != nil && watchCtx.Err() == nil {
providerLogger.Warn(watchCtx, "ai provider watch loop exited", slog.Error(err))
}
logger: gatewayLogger,
metrics: metrics,
providerMetrics: providerMetrics,
tracer: tracer,
})
defer func() {
watchCancel()
watchWG.Wait()
}()
mux := newGatewayMux(srv, srv.Ready, mw)
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
},
}
@@ -305,6 +235,203 @@ func (r *RootCmd) aiGatewayStart() *serpent.Command {
return cmd
}
type standaloneGatewayParams struct {
// Configuration.
bridgeConfig codersdk.AIBridgeConfig
coderURL string
httpAddress string
tlsCertFile string
tlsKeyFile string
// Runtime dependencies.
dialer aibridged.Dialer
pool aibridged.Pooler
// Observability.
// logger is the gateway-scoped logger; derived loggers (daemon,
// providers) are named under it.
logger slog.Logger
metrics *aibridge.Metrics
providerMetrics *aibridged.Metrics
tracer trace.Tracer
}
type standaloneGateway struct {
// Services.
daemon *aibridged.Server
httpServer *http.Server
reloader aibridged.ProviderReloader
// Configuration.
coderURL string
httpAddress string
tlsCertFile string
tlsKeyFile string
// State.
// providersLoaded is an initial-load latch. Reconnects refresh providers
// through the watch loop without resetting readiness.
providersLoaded atomic.Bool
// Observability.
logger slog.Logger
providerLogger slog.Logger
}
// runStandaloneGateway starts the aibridged daemon and serves the standalone
// AI Gateway. The daemon dials coderd asynchronously, so HTTP serving does not
// wait for the DRPC connection. It manages the daemon life cycle.
func runStandaloneGateway(ctx context.Context, params standaloneGatewayParams) error {
// The aibridged daemon must outlive ctx so in-flight HTTP requests
// retain their DRPC connection during graceful HTTP shutdown.
daemon, err := aibridged.New(context.Background(), params.pool, params.dialer, params.logger.Named("aibridged"), params.tracer)
if err != nil {
return xerrors.Errorf("start AI Gateway daemon: %w", err)
}
providerLogger := params.logger.Named("providers")
gateway := &standaloneGateway{
daemon: daemon,
reloader: agpl.NewPoolRPCReloader(params.pool, daemon.ClientContext, params.bridgeConfig, providerLogger, params.metrics, params.providerMetrics),
coderURL: params.coderURL,
httpAddress: params.httpAddress,
tlsCertFile: params.tlsCertFile,
tlsKeyFile: params.tlsKeyFile,
logger: params.logger,
providerLogger: providerLogger,
}
gateway.httpServer = &http.Server{
Handler: newGatewayMux(gateway.daemon, gateway.ready, gatewayMiddleware(params.bridgeConfig, params.tracer)),
ReadHeaderTimeout: time.Minute,
}
serveErr := gateway.serve(ctx)
var daemonShutdownErr error
if err := shutdownWithTimeout(daemon.Shutdown, daemonShutdownTimeout); err != nil {
daemonShutdownErr = xerrors.Errorf("shutdown AI Gateway daemon: %w", err)
}
return errors.Join(serveErr, daemonShutdownErr)
}
func (s *standaloneGateway) serve(ctx context.Context) error {
listener, err := net.Listen("tcp", s.httpAddress)
if err != nil {
return xerrors.Errorf("listen on %q: %w", s.httpAddress, err)
}
serveErr := make(chan error, 1)
var serveWG sync.WaitGroup
serveWG.Go(func() {
defer listener.Close()
if s.tlsCertFile != "" {
serveErr <- s.httpServer.ServeTLS(listener, s.tlsCertFile, s.tlsKeyFile)
return
}
serveErr <- s.httpServer.Serve(listener)
})
s.logger.Info(ctx, "standalone AI Gateway listening",
slog.F("address", listener.Addr().String()),
slog.F("coder_url", s.coderURL),
slog.F("tls", s.tlsCertFile != ""),
)
provReloadCtx, provReloadCancel := context.WithCancel(ctx)
provReloadDone := make(chan struct{})
go func() {
defer close(provReloadDone)
if err := s.loadProviders(provReloadCtx); err != nil {
if provReloadCtx.Err() == nil {
s.providerLogger.Error(provReloadCtx, "initial ai provider load stopped", slog.Error(err))
}
return
}
// WatchProviderReload reconnects internally and normally returns only when canceled.
err := aibridged.WatchProviderReload(provReloadCtx, s.daemon.ClientContext, s.reloader, s.providerLogger)
if err != nil && provReloadCtx.Err() == nil {
s.providerLogger.Error(provReloadCtx, "ai provider reload watch stopped", slog.Error(err))
}
}()
var runErr error
select {
case <-ctx.Done():
case <-s.daemon.Done():
// daemon uses context.Background() so no race with ctx.Done() is possible, ctx.Err() check is not needed.
runErr = xerrors.Errorf("AI Gateway daemon exited: %w", s.daemon.Err())
case <-provReloadDone:
if ctx.Err() == nil {
select {
// reload can exit due to daemon failure
// covering race with previous daemon.Done() case.
case <-s.daemon.Done():
runErr = xerrors.Errorf("AI Gateway daemon exited: %w", s.daemon.Err())
default:
runErr = xerrors.New("provider reload stopped unexpectedly")
}
}
case err := <-serveErr:
if err != nil && !errors.Is(err, http.ErrServerClosed) {
runErr = xerrors.Errorf("serve: %w", err)
}
}
s.logger.Info(ctx, "shutting down standalone AI Gateway")
provReloadCancel()
provReloadShutdownCtx, provReloadShutdownCancel := context.WithTimeout(context.Background(), providerReloadShutdownTimeout)
defer provReloadShutdownCancel()
var provReloadStopErr error
select {
case <-provReloadDone:
case <-provReloadShutdownCtx.Done():
provReloadStopErr = xerrors.Errorf("provider reload did not stop within %s, continuing gateway shutdown", providerReloadShutdownTimeout)
}
// Provider reload normally stops before HTTP draining so it cannot clear the
// bridge cache while requests are draining. If it does not stop within its
// timeout, continue with best-effort graceful HTTP shutdown.
// The daemon remains connected so in-flight requests retain their DRPC connection.
shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), httpShutdownTimeout)
defer shutdownCancel()
var httpShutdownErr error
if err := s.httpServer.Shutdown(shutdownCtx); err != nil {
httpShutdownErr = xerrors.Errorf("shutdown http server: %w", err)
if closeErr := s.httpServer.Close(); closeErr != nil {
httpShutdownErr = errors.Join(httpShutdownErr, xerrors.Errorf("force close http server: %w", closeErr))
}
}
serveWG.Wait()
return errors.Join(runErr, provReloadStopErr, httpShutdownErr)
}
// loadProviders retries the initial provider load until it succeeds or the
// context or daemon stops. A successful empty provider list completes the
// initial load. Subsequent changes are handled by the watch loop.
func (s *standaloneGateway) loadProviders(ctx context.Context) error {
for r := retry.New(50*time.Millisecond, 10*time.Second); r.Wait(ctx); {
if err := s.reloader.Reload(ctx); err != nil {
select {
case <-s.daemon.Done():
return err
default:
}
s.providerLogger.Warn(ctx, "failed to load ai providers, will retry", slog.Error(err))
continue
}
s.providersLoaded.Store(true)
s.providerLogger.Info(ctx, "loaded ai providers from coderd")
return nil
}
return context.Cause(ctx)
}
func (s *standaloneGateway) ready() bool {
return s.daemon.Ready() && s.providersLoaded.Load()
}
func gatewayMiddleware(cfg codersdk.AIBridgeConfig, tracer trace.Tracer) func(http.Handler) http.Handler {
mw := coderd.AIGatewayDataPlaneMiddleware(cfg)
// Tracing wraps outermost so rejected requests are still traced.
@@ -314,30 +441,6 @@ func gatewayMiddleware(cfg codersdk.AIBridgeConfig, tracer trace.Tracer) func(ht
}
}
func newGatewayMux(aibridgedHandler http.Handler, aibridgedReady func() bool, middleware func(http.Handler) http.Handler) *http.ServeMux {
mux := http.NewServeMux()
mux.Handle("/api/v2/aibridge/", middleware(http.StripPrefix("/api/v2/aibridge", aibridgedHandler)))
mux.Handle("/api/v2/ai-gateway/", middleware(http.StripPrefix("/api/v2/ai-gateway", aibridgedHandler)))
mux.Handle("/", middleware(aibridgedHandler))
// Health probes are registered without middleware.
mux.HandleFunc(healthzPath, func(w http.ResponseWriter, _ *http.Request) {
// healthz: returns 200 once the HTTP server is listening.
w.WriteHeader(http.StatusOK)
})
mux.HandleFunc(readyzPath, func(w http.ResponseWriter, _ *http.Request) {
// readyz: returns 200 only when the DRPC connection to coderd is established.
if aibridgedReady() {
w.WriteHeader(http.StatusOK)
return
}
w.WriteHeader(http.StatusServiceUnavailable)
})
return mux
}
// tracingMiddleware traces every request to the wrapped handler, unlike
// tracing.Middleware which only spans coderd's route patterns.
func tracingMiddleware(tracer trace.Tracer) func(http.Handler) http.Handler {
@@ -363,6 +466,31 @@ func tracingMiddleware(tracer trace.Tracer) func(http.Handler) http.Handler {
}
}
func newGatewayMux(aibridgedHandler http.Handler, aibridgedReady func() bool, middleware func(http.Handler) http.Handler) *http.ServeMux {
mux := http.NewServeMux()
mux.Handle("/api/v2/aibridge/", middleware(http.StripPrefix("/api/v2/aibridge", aibridgedHandler)))
mux.Handle("/api/v2/ai-gateway/", middleware(http.StripPrefix("/api/v2/ai-gateway", aibridgedHandler)))
mux.Handle("/", middleware(aibridgedHandler))
// Health probes are registered without middleware.
mux.HandleFunc(healthzPath, func(w http.ResponseWriter, _ *http.Request) {
// healthz: returns 200 once the HTTP server is listening.
w.WriteHeader(http.StatusOK)
})
mux.HandleFunc(readyzPath, func(w http.ResponseWriter, _ *http.Request) {
// readyz: returns 200 after the initial provider load while the
// DRPC connection to coderd remains active.
if aibridgedReady() {
w.WriteHeader(http.StatusOK)
return
}
w.WriteHeader(http.StatusServiceUnavailable)
})
return mux
}
// 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) {
@@ -381,33 +509,3 @@ func resolveAIGatewayKey(key string, keyFile string) (string, error) {
}
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.
//
// Subsequent provider changes are delivered by WatchProviderReload, started
// after this initial load returns.
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()
}
+487 -76
View File
@@ -4,39 +4,30 @@ package cli
import (
"context"
"errors"
"fmt"
"net"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"sync"
"sync/atomic"
"testing"
"github.com/stretchr/testify/require"
sdktrace "go.opentelemetry.io/otel/sdk/trace"
"golang.org/x/xerrors"
"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"
)
// 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.
@@ -52,88 +43,508 @@ func (r *failThenSucceedReloader) Reload(_ context.Context) error {
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{}
type failingReloader struct {
after func()
calls atomic.Int32
err error
}
func (r *alwaysFailReloader) Reload(context.Context) error {
func (r *failingReloader) 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()
type connectedDRPCConn struct {
drpc.Conn
closed chan struct{}
once sync.Once
}
// 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()
func (c *connectedDRPCConn) Close() error {
c.once.Do(func() {
close(c.closed)
})
return nil
}
func (c *connectedDRPCConn) Closed() <-chan struct{} {
return c.closed
}
type controlledShutdownPool struct {
*aibridged.CachedBridgePool
err error
release <-chan struct{}
started chan<- struct{}
}
func (p *controlledShutdownPool) Shutdown(ctx context.Context) error {
if p.started != nil {
p.started <- struct{}{}
}
if p.release != nil {
select {
case <-p.release:
case <-ctx.Done():
return errors.Join(ctx.Err(), p.err)
}
}
return errors.Join(p.CachedBridgePool.Shutdown(ctx), p.err)
}
type standaloneGatewayTestParams struct {
address string
params standaloneGatewayParams
pool *controlledShutdownPool
}
func newStandaloneGatewayTestParams(t *testing.T) *standaloneGatewayTestParams {
t.Helper()
reloader := &blockingReloader{started: make(chan struct{}, 1)}
logger := slog.Make()
tracer := sdktrace.NewTracerProvider().Tracer("test")
cachedPool, err := aibridged.NewCachedBridgePool(aibridged.DefaultPoolOptions, nil, logger, nil, tracer)
require.NoError(t, err)
t.Cleanup(func() {
require.NoError(t, shutdownWithTimeout(cachedPool.Shutdown, testutil.WaitShort))
})
done := make(chan error, 1)
go func() {
done <- loadProviders(runCtx, reloader, logger, nil)
}()
pool := &controlledShutdownPool{CachedBridgePool: cachedPool}
address := fmt.Sprintf("127.0.0.1:%d", testutil.RandomPort(t))
return &standaloneGatewayTestParams{
address: address,
params: standaloneGatewayParams{
httpAddress: address,
// Wait for the reload to be in-flight, then cancel as a signal would.
testutil.RequireReceive(testCtx, t, reloader.started)
cancel()
dialer: blockingStandaloneDaemonDialer,
pool: pool,
err := testutil.RequireReceive(testCtx, t, done)
require.ErrorIs(t, err, context.Canceled)
logger: logger,
tracer: tracer,
},
pool: pool,
}
}
// 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) {
func TestStandaloneGatewayLoadProviders(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)
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: "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: "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
},
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
},
wantErr: context.Canceled,
wantCalls: 1,
},
}
err := loadProviders(ctx, reloader, slog.Make(), aibridgedDone)
require.ErrorContains(t, err, errMsg)
require.Equal(t, int32(1), reloader.calls.Load())
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
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,
}
err := gateway.loadProviders(ctx)
if tc.wantErr == nil {
require.NoError(t, err)
} else {
require.ErrorIs(t, err, tc.wantErr)
}
require.Equal(t, tc.wantCalls, calls.Load())
require.Equal(t, tc.wantLoaded, gateway.providersLoaded.Load())
})
}
}
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()
}
}
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,
}
// The HTTP server is healthy before the daemon connects or providers load.
require.Equal(t, http.StatusOK, healthzStatus(t, gateway))
require.Equal(t, http.StatusServiceUnavailable, readyzStatus(t, gateway))
// 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.Equal(t, http.StatusOK, healthzStatus(t, gateway))
require.Equal(t, http.StatusServiceUnavailable, readyzStatus(t, gateway))
// The gateway becomes ready after the initial provider load completes.
require.NoError(t, gateway.loadProviders(ctx))
require.Equal(t, http.StatusOK, healthzStatus(t, gateway))
require.Equal(t, http.StatusOK, readyzStatus(t, gateway))
// 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.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.Equal(t, http.StatusOK, healthzStatus(t, gateway))
require.Equal(t, http.StatusOK, readyzStatus(t, gateway))
}
func healthzStatus(t *testing.T, gateway *standaloneGateway) int {
t.Helper()
return probeStatus(t, gateway, healthzPath)
}
func readyzStatus(t *testing.T, gateway *standaloneGateway) int {
t.Helper()
return probeStatus(t, gateway, readyzPath)
}
func probeStatus(t *testing.T, gateway *standaloneGateway, path string) int {
t.Helper()
rec := httptest.NewRecorder()
gateway.httpServer.Handler.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, path, nil))
return rec.Code
}
func TestAIGatewayStart_HealthBeforeProviders(t *testing.T) {
t.Parallel()
gatewayAddress := fmt.Sprintf("127.0.0.1:%d", testutil.RandomPort(t))
coderAddress := 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", "http://"+coderAddress,
"ai-gateway", "start",
"--key", "test-key",
"--http-address", gatewayAddress,
)
clitest.Start(t, inv.WithContext(testutil.Context(t, testutil.WaitShort)))
client := &http.Client{Timeout: testutil.WaitShort}
baseURL := "http://" + gatewayAddress
require.Eventually(t, func() bool {
req, err := http.NewRequestWithContext(t.Context(), 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.WaitShort, testutil.IntervalFast)
req, err := http.NewRequestWithContext(t.Context(), 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)
runDone := make(chan error, 1)
go func() {
runDone <- runStandaloneGateway(runCtx, test.params)
}()
requireListenerReady(t, test.address)
cancelRun()
require.NoError(t, testutil.RequireReceive(testCtx, t, runDone))
requireListenerAvailable(t, test.address, "HTTP listener must be closed before run returns")
}
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"})
}
err := runStandaloneGateway(testutil.Context(t, testutil.WaitShort), test.params)
require.ErrorContains(t, err, "AI Gateway daemon exited")
requireListenerAvailable(t, test.address, "HTTP listener must be closed before run returns")
}
func TestRunStandaloneGateway_HTTPStopsBeforeDaemonShutdown(t *testing.T) {
t.Parallel()
testCtx := testutil.Context(t, testutil.WaitShort)
test := newStandaloneGatewayTestParams(t)
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")
runDone := make(chan error, 1)
go func() {
runDone <- runStandaloneGateway(testCtx, test.params)
}()
testutil.RequireReceive(testCtx, t, shutdownStarted)
requireListenerAvailable(t, test.address, "HTTP listener must close before daemon shutdown")
close(shutdownRelease)
err := testutil.RequireReceive(testCtx, t, runDone)
require.ErrorContains(t, err, "serve:")
require.ErrorContains(t, err, "shutdown AI Gateway daemon:")
require.ErrorContains(t, err, shutdownErr.Error())
}
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", test.address)
require.NoError(t, err)
t.Cleanup(func() {
_ = listener.Close()
})
err = runStandaloneGateway(testutil.Context(t, testutil.WaitShort), test.params)
require.NoError(t, listener.Close())
require.ErrorContains(t, err, "listen on")
require.ErrorContains(t, err, "shutdown AI Gateway daemon:")
require.ErrorContains(t, err, shutdownErr.Error())
}
func TestStandaloneGatewayServe_ShutdownOrder(t *testing.T) {
t.Parallel()
// 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)
dialCtxCh := make(chan context.Context, 1)
dialer := func(ctx context.Context) (aibridged.DRPCClient, error) {
select {
case dialCtxCh <- ctx:
default:
}
<-ctx.Done()
return nil, ctx.Err()
}
daemon, err := aibridged.New(context.Background(), pool, dialer, logger, sdktrace.NewTracerProvider().Tracer("test"))
require.NoError(t, err)
httpAddress := fmt.Sprintf("127.0.0.1:%d", testutil.RandomPort(t))
reloader := &failThenSucceedReloader{}
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, _ *http.Request) {
select {
case handlerStarted <- struct{}{}:
default:
}
<-releaseHandler
w.WriteHeader(http.StatusNoContent)
}),
ReadHeaderTimeout: testutil.WaitShort,
},
httpAddress: httpAddress,
logger: logger,
providerLogger: logger,
reloader: reloader,
}
gateway.httpServer.RegisterOnShutdown(func() {
httpShutdownStarted <- struct{}{}
})
serveCtx, cancelServe := context.WithCancel(testCtx)
serveDone := make(chan error, 1)
go func() {
serveDone <- gateway.serve(serveCtx)
}()
dialCtx := testutil.RequireReceive(testCtx, t, dialCtxCh)
require.Eventually(t, gateway.providersLoaded.Load, testutil.WaitShort, testutil.IntervalFast)
requireListenerReady(t, httpAddress)
requestDone := make(chan error, 1)
go func() {
req, err := http.NewRequestWithContext(testCtx, http.MethodGet, "http://"+httpAddress, nil)
if err != nil {
requestDone <- err
return
}
resp, err := http.DefaultClient.Do(req)
if err == nil {
defer resp.Body.Close()
if resp.StatusCode != http.StatusNoContent {
err = xerrors.Errorf("unexpected status code: %d", resp.StatusCode)
}
}
requestDone <- err
}()
testutil.RequireReceive(testCtx, t, handlerStarted)
// Trigger shutdown after the initial load enters the provider watch loop.
cancelServe()
testutil.RequireReceive(testCtx, t, httpShutdownStarted)
select {
case <-dialCtx.Done():
t.Fatal("daemon context canceled before the in-flight HTTP request drained")
case err := <-serveDone:
t.Fatalf("server returned before the in-flight HTTP request drained: %v", err)
default:
}
// Expect provider reload to stop while HTTP draining keeps the daemon alive.
close(releaseHandler)
require.NoError(t, testutil.RequireReceive(testCtx, t, requestDone))
require.NoError(t, testutil.RequireReceive(testCtx, t, serveDone))
select {
case <-dialCtx.Done():
t.Fatal("daemon context canceled by serve")
case <-daemon.Done():
t.Fatal("daemon stopped before its runtime owner shut it down")
default:
}
// Expect the runtime owner to shut down the daemon after HTTP serving stops.
require.NoError(t, shutdownWithTimeout(daemon.Shutdown, daemonShutdownTimeout))
testutil.TryReceive(testCtx, t, dialCtx.Done())
testutil.TryReceive(testCtx, t, daemon.Done())
requireListenerAvailable(t, httpAddress, "HTTP listener must be closed before serve returns")
}
func requireListenerReady(t *testing.T, address string) {
t.Helper()
ctx := testutil.Context(t, testutil.WaitShort)
testutil.Eventually(ctx, t, func(ctx context.Context) bool {
conn, err := (&net.Dialer{}).DialContext(ctx, "tcp", address)
if err != nil {
return false
}
_ = conn.Close()
return true
}, testutil.IntervalFast)
}
func requireListenerAvailable(t *testing.T, address, message string) {
t.Helper()
listener, err := net.Listen("tcp", address)
require.NoError(t, err, message)
require.NoError(t, listener.Close())
}
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()
}
func TestResolveAIGatewayKey(t *testing.T) {
+6
View File
@@ -101,6 +101,12 @@ spec:
image: ghcr.io/coder/coder:v2.36.0
imagePullPolicy: IfNotPresent
lifecycle: {}
livenessProbe:
httpGet:
path: /healthz
port: http
scheme: HTTP
initialDelaySeconds: 0
name: coder
ports:
- containerPort: 4001
@@ -102,6 +102,12 @@ spec:
image: ghcr.io/coder/coder:v2.36.0
imagePullPolicy: IfNotPresent
lifecycle: {}
livenessProbe:
httpGet:
path: /healthz
port: http
scheme: HTTPS
initialDelaySeconds: 0
name: coder
ports:
- containerPort: 4001
@@ -102,6 +102,12 @@ spec:
image: ghcr.io/coder/coder:v2.36.0
imagePullPolicy: IfNotPresent
lifecycle: {}
livenessProbe:
httpGet:
path: /healthz
port: http
scheme: HTTP
initialDelaySeconds: 0
name: coder
ports:
- containerPort: 4001
+6
View File
@@ -100,6 +100,12 @@ spec:
image: ghcr.io/coder/coder:v2.36.0
imagePullPolicy: IfNotPresent
lifecycle: {}
livenessProbe:
httpGet:
path: /healthz
port: http
scheme: HTTP
initialDelaySeconds: 0
name: coder
ports:
- containerPort: 4001
+6 -6
View File
@@ -117,17 +117,15 @@ coder:
memory: 1Gi
# coder.startupProbe -- Startup probe configuration for the AI Gateway.
# Enable this with a failure threshold long enough for initial provider load
# before enabling the liveness probe.
startupProbe:
enabled: false
initialDelaySeconds: 0
# coder.livenessProbe -- Liveness probe configuration for the AI Gateway.
# Without this probe, Kubernetes does not restart a running but unresponsive
# Gateway. Enable and tune the startup probe before enabling this probe.
# This checks HTTP server responsiveness. Readiness reflects the established
# connection to coderd and initial provider configuration loaded.
livenessProbe:
enabled: false
enabled: true
initialDelaySeconds: 0
# coder.readinessProbe -- Readiness probe configuration for the AI Gateway.
@@ -190,7 +188,9 @@ aigateway:
name: ""
certKey: tls.crt
keyKey: tls.key
# This must exceed the application's 300-second shutdown timeout.
# Allow 5 seconds for provider reload shutdown, 300 seconds for HTTP
# draining, 5 seconds for daemon shutdown, 5 seconds for tracing shutdown,
# and 15 seconds of termination headroom.
terminationGracePeriodSeconds: 330
# Stable data-plane Service fronting the container's port 4001 listener.