diff --git a/cli/aibridged.go b/cli/aibridged.go index 75ded313c977c..a3316609e2a56 100644 --- a/cli/aibridged.go +++ b/cli/aibridged.go @@ -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, diff --git a/enterprise/cli/aigatewaystart.go b/enterprise/cli/aigatewaystart.go index 7faba2f9703d2..26b4ffcc774d3 100644 --- a/enterprise/cli/aigatewaystart.go +++ b/enterprise/cli/aigatewaystart.go @@ -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() - - // 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) - } - - 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)) - } - }) - 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, - } + return runStandaloneGateway(signalCtx, standaloneGatewayParams{ + bridgeConfig: vals.AI.BridgeConfig, + coderURL: serverURL.String(), + httpAddress: httpAddress, + tlsCertFile: tlsCertFile, + tlsKeyFile: tlsKeyFile, - serveErr := make(chan error, 1) - go func() { - if tlsCertFile != "" { - serveErr <- httpServer.ServeTLS(listener, tlsCertFile, tlsKeyFile) - } else { - serveErr <- httpServer.Serve(listener) - } - }() + dialer: aibridged.NewWebsocketDialer(serverURL, transport, resolvedKey), + pool: pool, - 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 + logger: gatewayLogger, + metrics: metrics, + providerMetrics: providerMetrics, + tracer: tracer, + }) }, } @@ -305,37 +235,210 @@ func (r *RootCmd) aiGatewayStart() *serpent.Command { return cmd } -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. - traced := tracingMiddleware(tracer) - return func(next http.Handler) http.Handler { - return traced(mw(next)) +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 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)) +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) + } - // 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) + 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) }) - 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) + 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 } - w.WriteHeader(http.StatusServiceUnavailable) - }) + // 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") - return mux + 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. + traced := tracingMiddleware(tracer) + return func(next http.Handler) http.Handler { + return traced(mw(next)) + } } // tracingMiddleware traces every request to the wrapped handler, unlike @@ -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() -} diff --git a/enterprise/cli/aigatewaystart_internal_test.go b/enterprise/cli/aigatewaystart_internal_test.go index cd0e97ede7665..4ebb14d59512f 100644 --- a/enterprise/cli/aigatewaystart_internal_test.go +++ b/enterprise/cli/aigatewaystart_internal_test.go @@ -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) { +type connectedDRPCConn struct { + drpc.Conn + closed chan struct{} + once sync.Once +} + +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() + + 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)) + }) + + pool := &controlledShutdownPool{CachedBridgePool: cachedPool} + address := fmt.Sprintf("127.0.0.1:%d", testutil.RandomPort(t)) + return &standaloneGatewayTestParams{ + address: address, + params: standaloneGatewayParams{ + httpAddress: address, + + dialer: blockingStandaloneDaemonDialer, + pool: pool, + + logger: logger, + tracer: tracer, + }, + pool: pool, + } +} + +func TestStandaloneGatewayLoadProviders(t *testing.T) { t.Parallel() - // testCtx bounds the test and drives the channel receives; runCtx is the - // context handed to loadProviders and is canceled to model a - // stop signal. They are distinct so the receives still work after the - // signal context is canceled. - testCtx := testutil.Context(t, testutil.WaitShort) - runCtx, cancel := context.WithCancel(testCtx) - defer cancel() + 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, + }, + } + + 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() - reloader := &blockingReloader{started: make(chan struct{}, 1)} + 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() - done := make(chan error, 1) + 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() { - done <- loadProviders(runCtx, reloader, logger, nil) + runDone <- runStandaloneGateway(runCtx, test.params) }() + requireListenerReady(t, test.address) + cancelRun() - // Wait for the reload to be in-flight, then cancel as a signal would. - testutil.RequireReceive(testCtx, t, reloader.started) - cancel() + require.NoError(t, testutil.RequireReceive(testCtx, t, runDone)) + requireListenerAvailable(t, test.address, "HTTP listener must be closed before run returns") +} - err := testutil.RequireReceive(testCtx, t, done) - require.ErrorIs(t, err, context.Canceled) +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") } -// 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 TestRunStandaloneGateway_HTTPStopsBeforeDaemonShutdown(t *testing.T) { t.Parallel() - ctx := testutil.Context(t, testutil.WaitShort) - reloader := &failThenSucceedReloader{failUntil: 2} + 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()) +} - require.NoError(t, loadProviders(ctx, reloader, slog.Make(), nil)) - require.GreaterOrEqual(t, reloader.calls.Load(), int32(3)) +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 TestLoadProviders_AIBridgedDoneStopsRetry(t *testing.T) { +func TestStandaloneGatewayServe_ShutdownOrder(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) + // 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) +} - err := loadProviders(ctx, reloader, slog.Make(), aibridgedDone) - require.ErrorContains(t, err, errMsg) - require.Equal(t, int32(1), reloader.calls.Load()) +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) { diff --git a/helm/ai-gateway/tests/testdata/default_values.golden b/helm/ai-gateway/tests/testdata/default_values.golden index cdee78607ffaf..66be5cd0566ec 100644 --- a/helm/ai-gateway/tests/testdata/default_values.golden +++ b/helm/ai-gateway/tests/testdata/default_values.golden @@ -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 diff --git a/helm/ai-gateway/tests/testdata/listener_tls_with_ingress.golden b/helm/ai-gateway/tests/testdata/listener_tls_with_ingress.golden index 75418126d3a9a..8b901086024fb 100644 --- a/helm/ai-gateway/tests/testdata/listener_tls_with_ingress.golden +++ b/helm/ai-gateway/tests/testdata/listener_tls_with_ingress.golden @@ -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 diff --git a/helm/ai-gateway/tests/testdata/networking_ai-gateway-test.golden b/helm/ai-gateway/tests/testdata/networking_ai-gateway-test.golden index 2374dbc0400ad..501ed02ad72e1 100644 --- a/helm/ai-gateway/tests/testdata/networking_ai-gateway-test.golden +++ b/helm/ai-gateway/tests/testdata/networking_ai-gateway-test.golden @@ -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 diff --git a/helm/ai-gateway/tests/testdata/nodeport.golden b/helm/ai-gateway/tests/testdata/nodeport.golden index 423e3b1d21fe0..e816359001617 100644 --- a/helm/ai-gateway/tests/testdata/nodeport.golden +++ b/helm/ai-gateway/tests/testdata/nodeport.golden @@ -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 diff --git a/helm/ai-gateway/values.yaml b/helm/ai-gateway/values.yaml index a3232d4566140..4d8b43d47441f 100644 --- a/helm/ai-gateway/values.yaml +++ b/helm/ai-gateway/values.yaml @@ -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.