From e8226a6d31ccc86276e42a4a8959f6951cf1a723 Mon Sep 17 00:00:00 2001 From: Goodness Date: Sat, 25 Jul 2026 14:19:58 +0100 Subject: [PATCH] test: verify api graceful shutdown ordering --- cmd/api/main.go | 112 +++++++++++++++++++++---- cmd/api/main_test.go | 195 +++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 291 insertions(+), 16 deletions(-) create mode 100644 cmd/api/main_test.go diff --git a/cmd/api/main.go b/cmd/api/main.go index ff883179..8d0dfb28 100644 --- a/cmd/api/main.go +++ b/cmd/api/main.go @@ -102,10 +102,6 @@ func main() { "max_conns", cfg.DBMaxConns, ) database = d - defer func() { - slog.Info("closing database connection") - database.Close() - }() if cfg.AutoMigrate { slog.Info("checking if migrations are needed", "step", "5", "action", "checking_migrations") @@ -123,7 +119,7 @@ func main() { slog.Info("migrations needed, running database migrations", "step", "5", "action", "running_database_migrations") // Use background context - migrations handle their own retries without timeouts allowIrreversible := cfg.IsDev() || os.Getenv("MIGRATE_ALLOW_IRREVERSIBLE") == "1" - err := migrate.Up(context.Background(), database.Pool, allowIrreversible) + err := migrate.Up(context.Background(), database.Pool, allowIrreversible) if err != nil { slog.Error("migration failed", "step", "5", "action", "migration_failed", "error", err, @@ -154,10 +150,6 @@ func main() { } slog.Info("nats connection successful", "step", "6.2", "action", "nats_connection_successful") eventBus = b - defer func() { - slog.Info("closing NATS connection") - eventBus.Close() - }() } else { slog.Info("nats skipped", "step", "6", "action", "nats_skipped", "reason", "NATS_URL not set") } @@ -239,17 +231,105 @@ func main() { ctx, cancel := context.WithTimeout(context.Background(), cfg.ShutdownTimeout) defer cancel() - if err := api.Shutdown(ctx, app); err != nil { + runner := gracefulShutdownRunner{deps: gracefulShutdownDeps{ + ShutdownHTTP: func(ctx context.Context) error { + return api.Shutdown(ctx, app) + }, + StopWorkers: stopWorkers, + WaitWorkers: func(ctx context.Context) error { + return shutdownwait.Wait(ctx, &workerWG) + }, + CloseBus: func(context.Context) error { + if eventBus == nil { + return nil + } + slog.Info("closing NATS connection") + eventBus.Close() + return nil + }, + CloseDB: func(context.Context) error { + if database == nil { + return nil + } + slog.Info("closing database connection") + database.Close() + return nil + }, + }} + result := runner.Run(ctx) + if result.WorkerErr != nil { + slog.Warn("worker shutdown exceeded deadline", "error", result.WorkerErr) + } + if result.Err != nil { slog.Error("graceful shutdown failed", - "error", err, - "error_type", fmt.Sprintf("%T", err), + "error", result.Err, + "error_type", fmt.Sprintf("%T", result.Err), ) os.Exit(1) } - stopWorkers() - if err := shutdownwait.Wait(ctx, &workerWG); err != nil { - slog.Warn("worker shutdown exceeded deadline", "error", err) - } slog.Info("shutdown complete") } + +type gracefulShutdownDeps struct { + ShutdownHTTP func(context.Context) error + StopWorkers func() + WaitWorkers func(context.Context) error + CloseBus func(context.Context) error + CloseDB func(context.Context) error +} + +type gracefulShutdownResult struct { + Err error + WorkerErr error +} + +type gracefulShutdownRunner struct { + deps gracefulShutdownDeps + once sync.Once + res gracefulShutdownResult +} + +func (r *gracefulShutdownRunner) Run(ctx context.Context) gracefulShutdownResult { + r.once.Do(func() { + r.res = runGracefulShutdown(ctx, r.deps) + }) + return r.res +} + +// runGracefulShutdown preserves the API process shutdown order. The HTTP +// listener stops first so new connections are rejected while in-flight requests +// can drain with the database still open. Workers are then canceled and waited +// on, the message bus is drained/closed, and the database pool is closed last. +func runGracefulShutdown(ctx context.Context, deps gracefulShutdownDeps) gracefulShutdownResult { + var shutdownErr error + + if deps.ShutdownHTTP != nil { + if err := deps.ShutdownHTTP(ctx); err != nil { + return gracefulShutdownResult{Err: fmt.Errorf("shutdown http listener: %w", err)} + } + } + + if deps.StopWorkers != nil { + deps.StopWorkers() + } + + var workerErr error + if deps.WaitWorkers != nil { + workerErr = deps.WaitWorkers(ctx) + } + + if deps.CloseBus != nil { + if err := deps.CloseBus(ctx); err != nil { + shutdownErr = errors.Join(shutdownErr, fmt.Errorf("close bus: %w", err)) + } + } + + if deps.CloseDB != nil { + if err := deps.CloseDB(ctx); err != nil { + shutdownErr = errors.Join(shutdownErr, fmt.Errorf("close database: %w", err)) + } + } + + return gracefulShutdownResult{Err: shutdownErr, WorkerErr: workerErr} +} diff --git a/cmd/api/main_test.go b/cmd/api/main_test.go new file mode 100644 index 00000000..afbd1306 --- /dev/null +++ b/cmd/api/main_test.go @@ -0,0 +1,195 @@ +package main + +import ( + "context" + "errors" + "reflect" + "strings" + "testing" +) + +func TestRunGracefulShutdownOrdersHTTPWorkersBusBeforeDB(t *testing.T) { + var order []string + httpStopped := false + busClosed := false + + result := runGracefulShutdown(context.Background(), gracefulShutdownDeps{ + ShutdownHTTP: func(context.Context) error { + order = append(order, "http") + httpStopped = true + return nil + }, + StopWorkers: func() { + order = append(order, "workers-stop") + }, + WaitWorkers: func(context.Context) error { + order = append(order, "workers-wait") + return nil + }, + CloseBus: func(context.Context) error { + if !httpStopped { + t.Fatal("bus closed before HTTP listener stopped accepting connections") + } + order = append(order, "bus") + busClosed = true + return nil + }, + CloseDB: func(context.Context) error { + if !httpStopped { + t.Fatal("database closed before HTTP listener stopped accepting connections") + } + if !busClosed { + t.Fatal("database closed before bus consumers drained/unsubscribed") + } + order = append(order, "db") + return nil + }, + }) + + if result.Err != nil { + t.Fatalf("runGracefulShutdown returned error: %v", result.Err) + } + if result.WorkerErr != nil { + t.Fatalf("runGracefulShutdown returned worker error: %v", result.WorkerErr) + } + + want := []string{"http", "workers-stop", "workers-wait", "bus", "db"} + if !reflect.DeepEqual(order, want) { + t.Fatalf("shutdown order = %v, want %v", order, want) + } +} + +func TestRunGracefulShutdownKeepsDBOpenUntilHTTPDrainCompletes(t *testing.T) { + requestDrained := false + + result := runGracefulShutdown(context.Background(), gracefulShutdownDeps{ + ShutdownHTTP: func(context.Context) error { + requestDrained = true + return nil + }, + CloseDB: func(context.Context) error { + if !requestDrained { + t.Fatal("database closed while HTTP shutdown was still draining in-flight work") + } + return nil + }, + }) + + if result.Err != nil { + t.Fatalf("runGracefulShutdown returned error: %v", result.Err) + } +} + +func TestRunGracefulShutdownDoesNotCloseDependenciesWhenHTTPShutdownFails(t *testing.T) { + httpErr := errors.New("listener still draining") + + result := runGracefulShutdown(context.Background(), gracefulShutdownDeps{ + ShutdownHTTP: func(context.Context) error { + return httpErr + }, + CloseBus: func(context.Context) error { + t.Fatal("bus closed after HTTP shutdown failed") + return nil + }, + CloseDB: func(context.Context) error { + t.Fatal("database closed after HTTP shutdown failed") + return nil + }, + }) + + if !errors.Is(result.Err, httpErr) { + t.Fatalf("shutdown error = %v, want wrapped HTTP error", result.Err) + } + if result.WorkerErr != nil { + t.Fatalf("worker error = %v, want nil", result.WorkerErr) + } +} + +func TestGracefulShutdownRunnerRunsOnceForRepeatedSignals(t *testing.T) { + var calls int + runner := gracefulShutdownRunner{ + deps: gracefulShutdownDeps{ + ShutdownHTTP: func(context.Context) error { + calls++ + return nil + }, + CloseBus: func(context.Context) error { + calls++ + return nil + }, + CloseDB: func(context.Context) error { + calls++ + return nil + }, + }, + } + + first := runner.Run(context.Background()) + second := runner.Run(context.Background()) + + if first.Err != nil || first.WorkerErr != nil { + t.Fatalf("first shutdown returned unexpected result: %+v", first) + } + if second.Err != nil || second.WorkerErr != nil { + t.Fatalf("second shutdown returned unexpected result: %+v", second) + } + if calls != 3 { + t.Fatalf("shutdown callbacks called %d times, want 3", calls) + } +} + +func TestRunGracefulShutdownReportsDependencyCloseErrors(t *testing.T) { + busErr := errors.New("bus drain failed") + dbErr := errors.New("database close failed") + var order []string + + result := runGracefulShutdown(context.Background(), gracefulShutdownDeps{ + CloseBus: func(context.Context) error { + order = append(order, "bus") + return busErr + }, + CloseDB: func(context.Context) error { + order = append(order, "db") + return dbErr + }, + }) + + if !errors.Is(result.Err, busErr) { + t.Fatalf("shutdown error %v does not wrap bus error", result.Err) + } + if !errors.Is(result.Err, dbErr) { + t.Fatalf("shutdown error %v does not wrap database error", result.Err) + } + if !strings.Contains(result.Err.Error(), "close bus") || !strings.Contains(result.Err.Error(), "close database") { + t.Fatalf("shutdown error lacks dependency context: %v", result.Err) + } + want := []string{"bus", "db"} + if !reflect.DeepEqual(order, want) { + t.Fatalf("close order = %v, want %v", order, want) + } +} + +func TestRunGracefulShutdownKeepsWorkerWaitTimeoutNonFatal(t *testing.T) { + waitErr := context.DeadlineExceeded + dbClosed := false + + result := runGracefulShutdown(context.Background(), gracefulShutdownDeps{ + WaitWorkers: func(context.Context) error { + return waitErr + }, + CloseDB: func(context.Context) error { + dbClosed = true + return nil + }, + }) + + if result.Err != nil { + t.Fatalf("worker wait timeout should not be treated as fatal shutdown error: %v", result.Err) + } + if !errors.Is(result.WorkerErr, waitErr) { + t.Fatalf("worker error = %v, want %v", result.WorkerErr, waitErr) + } + if !dbClosed { + t.Fatal("database close was skipped after worker wait timeout") + } +}