package main import ( "context" "errors" "go/ast" "go/parser" "go/token" "net/http" "reflect" "strings" "testing" "time" ) func TestRunRegistersSignalsBeforeStartingResources(t *testing.T) { fset := token.NewFileSet() file, err := parser.ParseFile(fset, "main.go", nil, 0) if err != nil { t.Fatalf("parse main.go: %v", err) } var runBody *ast.BlockStmt for _, declaration := range file.Decls { function, ok := declaration.(*ast.FuncDecl) if ok && function.Name.Name == "run" { runBody = function.Body break } } if runBody == nil { t.Fatal("main.go does not define run") } positions := make(map[string]token.Pos) ast.Inspect(runBody, func(node ast.Node) bool { call, ok := node.(*ast.CallExpr) if !ok { return true } selector, ok := call.Fun.(*ast.SelectorExpr) if !ok { return true } owner, ok := selector.X.(*ast.Ident) if !ok { return true } name := owner.Name + "." + selector.Sel.Name switch name { case "signal.NotifyContext", "db.NewPool", "queue.Start", "srv.ListenAndServe": positions[name] = call.Pos() } return true }) notifyPos, ok := positions["signal.NotifyContext"] if !ok { t.Fatal("run does not register for shutdown signals") } for _, start := range []string{"db.NewPool", "queue.Start", "srv.ListenAndServe"} { startPos, ok := positions[start] if !ok { t.Fatalf("run does not call %s", start) } if notifyPos >= startPos { t.Errorf("signal.NotifyContext at line %d must precede %s at line %d", fset.Position(notifyPos).Line, start, fset.Position(startPos).Line) } } } func TestWaitForShutdownReturnsListenerError(t *testing.T) { listenErr := errors.New("listen failed") serverErr := make(chan error, 1) serverErr <- listenErr err := waitForShutdown(make(chan struct{}), serverErr) if !errors.Is(err, listenErr) || !strings.Contains(err.Error(), "server") { t.Fatalf("waitForShutdown() error = %v, want wrapped listener error", err) } } func TestWaitForShutdownAcceptsSignalAndServerClosed(t *testing.T) { t.Run("signal", func(t *testing.T) { stop := make(chan struct{}) close(stop) if err := waitForShutdown(stop, make(chan error)); err != nil { t.Fatalf("waitForShutdown() error = %v, want nil", err) } }) t.Run("server closed", func(t *testing.T) { serverErr := make(chan error, 1) serverErr <- http.ErrServerClosed if err := waitForShutdown(make(chan struct{}), serverErr); err != nil { t.Fatalf("waitForShutdown() error = %v, want nil", err) } }) } func TestShutdownAllContinuesAfterTimeoutWithSharedDeadline(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), time.Minute) defer cancel() wantDeadline, _ := ctx.Deadline() var order []string var contexts []context.Context rootCtx, rootCancel := context.WithCancel(context.Background()) steps := []shutdownStep{ {name: "http", run: recordShutdown(&order, &contexts, "http", nil)}, {name: "upgraded", run: recordShutdown(&order, &contexts, "upgraded", nil)}, {name: "queue", run: recordShutdown(&order, &contexts, "queue", context.DeadlineExceeded)}, {name: "db", run: recordShutdown(&order, &contexts, "db", nil)}, } err := shutdownAll(ctx, rootCancel, steps...) if !errors.Is(err, context.DeadlineExceeded) { t.Fatalf("shutdownAll() error = %v, want deadline exceeded", err) } if want := []string{"http", "upgraded", "queue", "db", "root"}; !reflect.DeepEqual(orderWithRoot(order, rootCtx), want) { t.Fatalf("shutdown order = %v, want %v", orderWithRoot(order, rootCtx), want) } for i, gotCtx := range contexts { gotDeadline, ok := gotCtx.Deadline() if !ok || !gotDeadline.Equal(wantDeadline) { t.Errorf("step %d deadline = %v, %v; want %v, true", i, gotDeadline, ok, wantDeadline) } } } func TestShutdownAllContinuesWhenStepIgnoresDeadline(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), 20*time.Millisecond) defer cancel() blocked := make(chan struct{}) defer close(blocked) cleanupStarted := make(chan struct{}) rootCtx, rootCancel := context.WithCancel(context.Background()) result := make(chan error, 1) go func() { result <- shutdownAll(ctx, rootCancel, shutdownStep{name: "blocked", run: func(context.Context) error { <-blocked return nil }}, shutdownStep{name: "cleanup", run: func(context.Context) error { close(cleanupStarted) return nil }}, ) }() select { case err := <-result: if !errors.Is(err, context.DeadlineExceeded) { t.Fatalf("shutdownAll() error = %v, want deadline exceeded", err) } case <-time.After(250 * time.Millisecond): t.Fatal("shutdownAll did not enforce the shared deadline") } select { case <-cleanupStarted: default: t.Fatal("cleanup after blocked step was not started") } select { case <-rootCtx.Done(): default: t.Fatal("root context was not canceled") } } func recordShutdown(order *[]string, contexts *[]context.Context, name string, err error) func(context.Context) error { return func(ctx context.Context) error { *order = append(*order, name) *contexts = append(*contexts, ctx) return err } } func orderWithRoot(order []string, rootCtx context.Context) []string { got := append([]string(nil), order...) select { case <-rootCtx.Done(): got = append(got, "root") default: } return got }