diff --git a/app/app.go b/app/app.go index 4e1c7da..d13b646 100644 --- a/app/app.go +++ b/app/app.go @@ -4,9 +4,12 @@ import ( "context" "fmt" "log/slog" + "os" + "os/signal" + "syscall" + "github.com/draincloud/callpack/safegroup" "github.com/draincloud/logger" - "golang.org/x/sync/errgroup" ) type Runnable interface { @@ -15,39 +18,71 @@ type Runnable interface { type App struct { name string + strategy Strategy runnables []Runnable } -func NewApp( +type Strategy string + +const ( + // If one exits - all runnables will be killed with cancel + StrategyOneForAll Strategy = "one_for_all" + // If one exits WITHOUT error - there will be no context cancel. Error will still cause context cancellation. + StrategyOneForOne Strategy = "one_for_one" +) + +func New( name string, + strategy Strategy, runnables ...Runnable, ) *App { return &App{ name: name, + strategy: strategy, runnables: runnables, } } func (a *App) Run(ctx context.Context) error { ctx = logger.WithAttrs(ctx, slog.String("app", a.name)) - logger.Warn(ctx, "[App][Run] sstarting app") + logger.Warn(ctx, "[App][Run] starting app") - eg, egCtx := errgroup.WithContext(ctx) - - runCtx, cancel := context.WithCancel(egCtx) + ctx, cancel := signal.NotifyContext(ctx, os.Interrupt, syscall.SIGTERM) defer cancel() + errChan := make(chan error, 1) + defer close(errChan) + + stopChan := make(chan struct{}, 1) + defer close(stopChan) + + eg, egCtx := safegroup.WithContext(ctx) + + runCtx, runCancel := context.WithCancel(egCtx) + defer runCancel() + for _, r := range a.runnables { eg.Go(func() error { - defer cancel() - + if a.strategy == StrategyOneForAll { + defer runCancel() + } return r.Run(runCtx) }) } - if err := eg.Wait(); err != nil { - return fmt.Errorf("[app][Run] %s: %w", a.name, err) - } + go func() { + defer cancel() + if err := eg.Wait(); err != nil { + errChan <- fmt.Errorf("[app][Run] %s: %w", a.name, err) + return + } + stopChan <- struct{}{} + }() - return nil + select { + case err := <-errChan: + return err + case <-stopChan: + return nil + } } diff --git a/app/go.mod b/app/go.mod index b5920e3..f71dbc5 100644 --- a/app/go.mod +++ b/app/go.mod @@ -3,8 +3,8 @@ module github.com/draincloud/callpack/app go 1.26.3 require ( + github.com/draincloud/callpack/safegroup v0.1.0 github.com/draincloud/logger v0.0.5 - golang.org/x/sync v0.22.0 ) require ( diff --git a/app/go.sum b/app/go.sum index 3fcd09b..1fb01d8 100644 --- a/app/go.sum +++ b/app/go.sum @@ -1,3 +1,5 @@ +github.com/draincloud/callpack/safegroup v0.1.0 h1:iVFCnKYEBqZ3iSRLCwZhymkG+EZwmRQ9K4H3QkASop0= +github.com/draincloud/callpack/safegroup v0.1.0/go.mod h1:+WMZJmBlifqYU12NM4YIivzqCUbKfcRnIZ4AUeUwjlk= github.com/draincloud/logger v0.0.5 h1:4saeda/sm5E6S0lnfwCIdgmqZYiIGuTIUN0HpIlm3js= github.com/draincloud/logger v0.0.5/go.mod h1:Z/GP5qHAC+MrhTMs5FckdQahiP2MjugaeQxjaMMaqWY= github.com/fatih/color v1.18.0 h1:S8gINlzdQ840/4pfAwic/ZE0djQEH3wM94VfqLTZcOM= @@ -7,8 +9,6 @@ github.com/mattn/go-colorable v0.1.13/go.mod h1:7S9/ev0klgBDR4GtXTXX8a3vIGJpMovk github.com/mattn/go-isatty v0.0.16/go.mod h1:kYGgaQfpe5nmfYZH+SKPsOc2e4SrIfOl2e/yFXSvRLM= github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY= github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= -golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek= -golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= golang.org/x/sys v0.0.0-20220811171246-fbc7d0a398ab/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.25.0 h1:r+8e+loiHxRqhXVl6ML1nO3l1+oFoWbnlu2Ehimmi34= diff --git a/app/run_test.go b/app/run_test.go index 3db236c..2fa6856 100644 --- a/app/run_test.go +++ b/app/run_test.go @@ -3,6 +3,7 @@ package app_test import ( "context" "errors" + "sync" "testing" "time" @@ -23,23 +24,97 @@ func TestRunReturnsWhenRunnableExitsCleanly(t *testing.T) { }) oneShot := runnableFunc(func(context.Context) error { return nil }) - if err := runWithin(t, time.Second, app.NewApp("test", blocked, oneShot)); err != nil { + if err := runWithin(t, time.Second, app.New("test", app.StrategyOneForAll, blocked, oneShot)); err != nil { t.Fatalf("Run: %v", err) } } +func TestRunReturnsWhenRunnableOneForOne(t *testing.T) { + t.Parallel() + + blockedStopped := make(chan struct{}) + blocked := runnableFunc(func(ctx context.Context) error { + <-ctx.Done() + close(blockedStopped) + + return nil + }) + + oneCh := make(chan struct{}, 1) + oneShot := runnableFunc(func(context.Context) error { + oneCh <- struct{}{} + + return nil + }) + + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + + go func() { + _ = app.New("test", app.StrategyOneForOne, blocked, oneShot).Run(ctx) + }() + + select { + case <-oneCh: + case <-time.After(time.Second): + t.Fatal("oneShot did not run") + } + + // A clean exit under one-for-one must leave the other runnables alone. + select { + case <-blockedStopped: + t.Fatal("blocked app stopped") + case <-time.After(100 * time.Millisecond): + } +} + func TestRunWrapsRunnableError(t *testing.T) { t.Parallel() errBoom := errors.New("boom") failing := runnableFunc(func(context.Context) error { return errBoom }) - err := runWithin(t, time.Second, app.NewApp("test", failing)) + err := runWithin(t, time.Second, app.New("test", app.StrategyOneForAll, failing)) if !errors.Is(err, errBoom) { t.Fatalf("Run: got %v, want %v", err, errBoom) } } +func TestRunReturnsWhenRootCtxIsCanceledOneForOne(t *testing.T) { + t.Parallel() + + blocked := runnableFunc(func(ctx context.Context) error { + return nil + }) + + oneCh := make(chan struct{}, 1) + defer close(oneCh) + + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + + ticker := time.NewTicker(time.Second) + defer ticker.Stop() + + m := sync.Mutex{} + stopped := false + go func() { + _ = app.New("test", app.StrategyOneForOne, blocked).Run(ctx) + m.Lock() + stopped = true + m.Unlock() + }() + cancel() + + <-ticker.C + + m.Lock() + defer m.Unlock() + if !stopped { + t.Fatal("should be stopped but it is not") + } +} + func runWithin(t *testing.T, d time.Duration, a *app.App) error { t.Helper()