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
12 changes: 11 additions & 1 deletion script/engine.go
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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++
Expand Down
78 changes: 78 additions & 0 deletions script_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,8 +7,10 @@ import (
"bufio"
"bytes"
"context"
"fmt"
"strings"
"testing"
"time"

"github.com/cilium/hive"
"github.com/cilium/hive/cell"
Expand Down Expand Up @@ -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)
})
}
}
Loading