Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
112 changes: 96 additions & 16 deletions cmd/api/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand All @@ -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,
Expand Down Expand Up @@ -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")
}
Expand Down Expand Up @@ -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}
}
195 changes: 195 additions & 0 deletions cmd/api/main_test.go
Original file line number Diff line number Diff line change
@@ -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")
}
}