chore: fix port availability check flake from aigatewaystart_internal… (#27784)

Extends state tracking in `standaloneGateway` which is used in tests to
simplify checks.
Flaky `requireListenerAvailable` was removed.
This commit is contained in:
Paweł Banaszewski
2026-08-04 11:33:34 +02:00
committed by GitHub
parent 404bb2f663
commit a2287d6739
2 changed files with 61 additions and 56 deletions
+32 -11
View File
@@ -271,6 +271,9 @@ type standaloneGateway struct {
tlsKeyFile string
// State.
drpcClosed atomic.Bool
providerRefreshStopped atomic.Bool
listenerClosed atomic.Bool
// providersLoaded is an initial-load latch. Reconnects refresh providers
// through the watch loop without resetting readiness.
providersLoaded atomic.Bool
@@ -280,15 +283,23 @@ type standaloneGateway struct {
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.
// runStandaloneGateway starts the aibridged daemon, serves the standalone
// AI Gateway and manages the daemon life cycle. The daemon dials coderd
// asynchronously. HTTP serving does not wait for the DRPC connection.
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.
gateway, err := newStandaloneGateway(params)
if err != nil {
return err
}
return gateway.run(ctx)
}
func newStandaloneGateway(params standaloneGatewayParams) (*standaloneGateway, error) {
// The aibridged daemon must outlive the serving context 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)
return nil, xerrors.Errorf("start AI Gateway daemon: %w", err)
}
providerLogger := params.logger.Named("providers")
@@ -308,15 +319,23 @@ func runStandaloneGateway(ctx context.Context, params standaloneGatewayParams) e
Handler: newGatewayMux(gateway.daemon, gateway.ready, gatewayMiddleware(params.bridgeConfig, params.tracer)),
ReadHeaderTimeout: time.Minute,
}
return gateway, nil
}
serveErr := gateway.serve(ctx)
var daemonShutdownErr error
if err := shutdownWithTimeout(daemon.Shutdown, daemonShutdownTimeout); err != nil {
daemonShutdownErr = xerrors.Errorf("shutdown AI Gateway daemon: %w", err)
}
func (s *standaloneGateway) run(ctx context.Context) error {
serveErr := s.serve(ctx)
daemonShutdownErr := s.shutdownDaemon()
return errors.Join(serveErr, daemonShutdownErr)
}
func (s *standaloneGateway) shutdownDaemon() error {
if err := shutdownWithTimeout(s.daemon.Shutdown, daemonShutdownTimeout); err != nil {
return xerrors.Errorf("shutdown AI Gateway daemon: %w", err)
}
s.drpcClosed.Store(true)
return nil
}
func (s *standaloneGateway) serve(ctx context.Context) error {
listener, err := net.Listen("tcp", s.httpAddress)
if err != nil {
@@ -388,6 +407,7 @@ func (s *standaloneGateway) serve(ctx context.Context) error {
var provReloadStopErr error
select {
case <-provReloadDone:
s.providerRefreshStopped.Store(true)
case <-provReloadShutdownCtx.Done():
provReloadStopErr = xerrors.Errorf("provider reload did not stop within %s, continuing gateway shutdown", providerReloadShutdownTimeout)
}
@@ -406,6 +426,7 @@ func (s *standaloneGateway) serve(ctx context.Context) error {
}
}
serveWG.Wait()
s.listenerClosed.Store(true)
return errors.Join(runErr, provReloadStopErr, httpShutdownErr)
}
+29 -45
View File
@@ -328,15 +328,19 @@ func TestRunStandaloneGateway_ContextCanceled(t *testing.T) {
defer cancelRun()
test := newStandaloneGatewayTestParams(t)
gateway, err := newStandaloneGateway(test.params)
require.NoError(t, err)
runDone := make(chan error, 1)
go func() {
runDone <- runStandaloneGateway(runCtx, test.params)
runDone <- gateway.run(runCtx)
}()
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")
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")
require.True(t, gateway.listenerClosed.Load(), "HTTP listener must be closed before run returns")
}
func TestRunStandaloneGateway_DaemonExited(t *testing.T) {
@@ -347,9 +351,13 @@ func TestRunStandaloneGateway_DaemonExited(t *testing.T) {
return nil, codersdk.NewError(http.StatusUnauthorized, codersdk.Response{Message: "invalid gateway key"})
}
err := runStandaloneGateway(testutil.Context(t, testutil.WaitShort), test.params)
gateway, err := newStandaloneGateway(test.params)
require.NoError(t, err)
err = gateway.run(testutil.Context(t, testutil.WaitShort))
require.ErrorContains(t, err, "AI Gateway daemon exited")
requireListenerAvailable(t, test.address, "HTTP listener must be closed before run returns")
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")
require.True(t, gateway.listenerClosed.Load(), "HTTP listener must be closed before run returns")
}
func TestRunStandaloneGateway_HTTPStopsBeforeDaemonShutdown(t *testing.T) {
@@ -366,15 +374,20 @@ func TestRunStandaloneGateway_HTTPStopsBeforeDaemonShutdown(t *testing.T) {
test.params.tlsCertFile = filepath.Join(t.TempDir(), "missing.crt")
test.params.tlsKeyFile = filepath.Join(t.TempDir(), "missing.key")
gateway, err := newStandaloneGateway(test.params)
require.NoError(t, err)
runDone := make(chan error, 1)
go func() {
runDone <- runStandaloneGateway(testCtx, test.params)
runDone <- gateway.run(testCtx)
}()
testutil.RequireReceive(testCtx, t, shutdownStarted)
requireListenerAvailable(t, test.address, "HTTP listener must close before daemon shutdown")
require.True(t, gateway.providerRefreshStopped.Load(), "provider refresh must stop before daemon shutdown")
require.True(t, gateway.listenerClosed.Load(), "HTTP listener must close before daemon shutdown")
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:")
require.ErrorContains(t, err, shutdownErr.Error())
@@ -408,16 +421,7 @@ func TestStandaloneGatewayServe_ShutdownOrder(t *testing.T) {
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"))
daemon, err := aibridged.New(context.Background(), pool, blockingStandaloneDaemonDialer, logger, sdktrace.NewTracerProvider().Tracer("test"))
require.NoError(t, err)
httpAddress := fmt.Sprintf("127.0.0.1:%d", testutil.RandomPort(t))
@@ -453,7 +457,6 @@ func TestStandaloneGatewayServe_ShutdownOrder(t *testing.T) {
serveDone <- gateway.serve(serveCtx)
}()
dialCtx := testutil.RequireReceive(testCtx, t, dialCtxCh)
require.Eventually(t, gateway.providersLoaded.Load, testutil.WaitShort, testutil.IntervalFast)
requireListenerReady(t, httpAddress)
@@ -478,32 +481,21 @@ func TestStandaloneGatewayServe_ShutdownOrder(t *testing.T) {
// 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:
}
require.True(t, gateway.providerRefreshStopped.Load(), "provider refresh must stop before HTTP draining")
require.False(t, gateway.listenerClosed.Load(), "HTTP listener must remain open while requests drain")
require.False(t, gateway.drpcClosed.Load(), "DRPC connection must remain open while requests drain")
// 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:
}
require.True(t, gateway.providerRefreshStopped.Load(), "provider refresh must stop before serve returns")
require.True(t, gateway.listenerClosed.Load(), "HTTP listener must be closed before serve returns")
require.False(t, gateway.drpcClosed.Load(), "DRPC connection must remain open until its runtime owner shuts it down")
// 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")
require.NoError(t, gateway.shutdownDaemon())
require.True(t, gateway.drpcClosed.Load(), "DRPC connection must close during daemon shutdown")
}
func requireListenerReady(t *testing.T, address string) {
@@ -520,14 +512,6 @@ func requireListenerReady(t *testing.T, address string) {
}, 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()