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
6 changes: 4 additions & 2 deletions .github/workflows/test.yml
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,8 @@ on:
pull_request:
branches:
- main
permissions:
contents: read
jobs:
reuse:
uses: bsm/misc/.github/workflows/test-go.yaml@main
go:
uses: bsm/misc/.github/workflows/test-go.yml@main
Comment thread
github-advanced-security[bot] marked this conversation as resolved.
Fixed
3 changes: 2 additions & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@ Wait for servers to terminate gracefully.
```go
import (
"context"
"errors"
"log"
"net/http"
"time"
Expand All @@ -26,7 +27,7 @@ func main() {
// Wait for either SIGINT/SIGTERM or ListenAndServe to exit.
// Handle errors.
err := shutdown.Wait(srv.ListenAndServe)
if err != nil && err != http.ErrServerClosed {
if err != nil && !errors.Is(err, http.ErrServerClosed) {
log.Fatalln("Server error", err)
}

Expand Down
3 changes: 2 additions & 1 deletion README.md.tpl
Original file line number Diff line number Diff line change
Expand Up @@ -10,12 +10,13 @@ Wait for servers to terminate gracefully.
```go
import (
"context"
"errors"
"log"
"net/http"
"time"

"github.com/bsm/shutdown"
)

func main() {{ "Example" | code }}
func main() {{ "ExampleWait" | code }}
```
5 changes: 4 additions & 1 deletion graceful.go
Original file line number Diff line number Diff line change
Expand Up @@ -36,7 +36,10 @@ func GracefulContext(ctx context.Context, start StartFunc, shutdown ShutdownFunc
return err
}

timeout, cancel := context.WithTimeout(context.Background(), DefaultShutdownTimeout)
// Detach the parent's cancellation/deadline so shutdown always gets the full
// timeout (the parent may already be cancelled, e.g. by the signal that
// triggered the shutdown), while still preserving any values it carries.
timeout, cancel := context.WithTimeout(context.WithoutCancel(ctx), DefaultShutdownTimeout)
defer cancel()

if err := shutdown(timeout); err != nil {
Expand Down
55 changes: 55 additions & 0 deletions graceful_test.go
Original file line number Diff line number Diff line change
@@ -1,8 +1,11 @@
package shutdown_test

import (
"context"
"errors"
"log"
"net/http"
"testing"

"github.com/bsm/shutdown"
)
Expand All @@ -17,3 +20,55 @@ func ExampleGraceful() {
log.Fatalln("Server error", err)
}
}

func TestGracefulContext_start_error(t *testing.T) {
boom := errors.New("boom")

// Unexpected start errors are surfaced and shutdown is not invoked.
called := false
err := shutdown.GracefulContext(context.Background(),
func() error { return boom },
func(context.Context) error { called = true; return nil },
)
if !errors.Is(err, boom) {
t.Fatalf("expected boom, got %v", err)
}
if called {
t.Fatal("shutdown should not run when start fails unexpectedly")
}

// Expected start errors are swallowed.
err = shutdown.GracefulContext(context.Background(),
func() error { return boom },
func(context.Context) error { return nil },
boom,
)
if err != nil {
t.Fatalf("expected nil, got %v", err)
}
}

type ctxKey string

func TestGracefulContext_shutdown_detached(t *testing.T) {
// A cancelled parent triggers shutdown; the shutdown context must remain
// live (fresh timeout) yet still carry the parent's values.
parent, cancel := context.WithCancel(context.WithValue(context.Background(), ctxKey("k"), "v"))
cancel()

err := shutdown.GracefulContext(parent,
func() error { <-parent.Done(); return nil },
func(ctx context.Context) error {
if ctx.Err() != nil {
t.Errorf("shutdown context already cancelled: %v", ctx.Err())
}
if got := ctx.Value(ctxKey("k")); got != "v" {
t.Errorf("parent value not preserved, got %v", got)
}
return nil
},
)
if err != nil {
t.Fatalf("expected nil, got %v", err)
}
}
15 changes: 8 additions & 7 deletions shutdown_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@ package shutdown_test

import (
"context"
"fmt"
"errors"
"log"
"net/http"
"testing"
Expand All @@ -20,7 +20,7 @@ func ExampleWait() {
// Wait for either SIGINT/SIGTERM or ListenAndServe to exit.
// Handle errors.
err := shutdown.Wait(srv.ListenAndServe)
if err != nil && err != http.ErrServerClosed {
if err != nil && !errors.Is(err, http.ErrServerClosed) {
log.Fatalln("Server error", err)
}

Expand All @@ -35,11 +35,12 @@ func ExampleWait() {
}

func TestWait_fails_immediately(t *testing.T) {
err := shutdown.Wait(func() error { return fmt.Errorf("doh!") })
sentinel := errors.New("doh!")
err := shutdown.Wait(func() error { return sentinel })
if err == nil {
t.Fatalf("expected error, got nil")
} else if err.Error() != "doh!" {
t.Fatalf("expected speficic error, got %v", err)
} else if !errors.Is(err, sentinel) {
t.Fatalf("expected specific error, got %v", err)
}
}

Expand All @@ -56,7 +57,7 @@ func TestWaitContext_nil_callback(t *testing.T) {
}
if err := ctx.Err(); err == nil {
t.Fatalf("expected error, got nil")
} else if err != context.Canceled {
t.Fatalf("expected speficic error, got %v", err)
} else if !errors.Is(err, context.Canceled) {
t.Fatalf("expected specific error, got %v", err)
}
}
Loading