diff --git a/script/engine.go b/script/engine.go index a1f2d19..bdd579a 100644 --- a/script/engine.go +++ b/script/engine.go @@ -86,6 +86,11 @@ type Engine struct { // MaxRetryInterval is the maximum time to wait before retrying. MaxRetryInterval time.Duration + + // MaxRetries is the maximum number of times (excluding the first one) a + // retrying command marked with '*' is executed before giving up. If zero, + // the number of retries is bound only by the context. + MaxRetries uint } // NewEngine returns an Engine configured with a basic set of commands and conditions. @@ -360,12 +365,17 @@ func (e *Engine) Execute(s *State, file string, script *bufio.Reader, log io.Wri // Command wants retries. Retry the whole section backoff := exponentialBackoff{max: maxRetryInterval, interval: retryInterval} for err != nil { + if e.MaxRetries != 0 && s.RetryCount >= int(e.MaxRetries) { + s.RetryCount = 0 + return lineErr(err) + } + retryDuration := backoff.get() fmt.Fprintf(log, "(command %q failed, retrying in %s...)\n", line, retryDuration) select { case <-s.Context().Done(): s.RetryCount = 0 - return lineErr(s.Context().Err()) + return lineErr(err) case <-time.After(retryDuration): } s.RetryCount++ diff --git a/script_test.go b/script_test.go index 0838b58..d46ae51 100644 --- a/script_test.go +++ b/script_test.go @@ -7,8 +7,10 @@ import ( "bufio" "bytes" "context" + "fmt" "strings" "testing" + "time" "github.com/cilium/hive" "github.com/cilium/hive/cell" @@ -73,3 +75,79 @@ hive/stop expected := `> hive/start.*> example1.*hello1.*> example2.*hello2.*> hive/stop` require.Regexp(t, expected, strings.ReplaceAll(stdout.String(), "\n", " ")) } + +func TestScriptCommandRetries(t *testing.T) { + var tests = []struct { + name string + succeedAfter uint + cancelAfter uint + maxRetries uint + assert require.ErrorAssertionFunc + }{ + { + name: "max two retries, succeed after two", + succeedAfter: 2, + maxRetries: 2, + assert: require.NoError, + }, + { + name: "max two retries, succeed after three", + succeedAfter: 3, + maxRetries: 2, + assert: func(tt require.TestingT, err error, args ...any) { + require.ErrorContains(t, err, "expected to succeed after 3 times, current: 2", args...) + }, + }, + { + name: "no limit, context cancellation only", + succeedAfter: 1000, + cancelAfter: 10, + assert: func(tt require.TestingT, err error, args ...any) { + require.ErrorContains(t, err, "expected to succeed after 1000 times, current: 10", args...) + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + var ( + counter uint + engine = script.Engine{ + Cmds: map[string]script.Cmd{ + "test": script.Command( + script.CmdUsage{}, + func(s *script.State, args ...string) (script.WaitFunc, error) { + defer func() { counter++ }() + + s.Logf("test command called %d times", counter) + if tt.cancelAfter != 0 && counter == tt.cancelAfter { + cancel() + } + + if counter != tt.succeedAfter { + return nil, fmt.Errorf("expected to succeed after %d times, current: %d", tt.succeedAfter, counter) + } + + return nil, nil + }, + ), + }, + + RetryInterval: 10 * time.Millisecond, + MaxRetryInterval: 10 * time.Millisecond, + MaxRetries: tt.maxRetries, + } + ) + + s, err := script.NewState(ctx, t.TempDir(), nil) + require.NoError(t, err, "NewState") + + var stdout bytes.Buffer + err = engine.Execute(s, "", bufio.NewReader(strings.NewReader("* test")), &stdout) + tt.assert(t, err) + }) + } +}