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