diff --git a/cmd/tap/firehose.go b/cmd/tap/firehose.go index 02f82f6d..4eedb87b 100644 --- a/cmd/tap/firehose.go +++ b/cmd/tap/firehose.go @@ -368,21 +368,16 @@ func (fp *FirehoseProcessor) saveCursor(ctx context.Context) error { // RunCursorSaver periodically saves the firehose cursor to the database. func (fp *FirehoseProcessor) RunCursorSaver(ctx context.Context) { - ticker := time.NewTicker(fp.cursorSaveInterval) - defer ticker.Stop() - - for { - select { - case <-ctx.Done(): - if err := fp.saveCursor(ctx); err != nil { - fp.logger.Error("failed to save cursor on shutdown", "error", err, "relayUrl", fp.relayUrl) - } - return - case <-ticker.C: - if err := fp.saveCursor(ctx); err != nil { - fp.logger.Error("failed to save cursor", "error", err, "relayUrl", fp.relayUrl) - } + runPeriodically(ctx, fp.cursorSaveInterval, func(ctx context.Context) error { + if err := fp.saveCursor(ctx); err != nil { + fp.logger.Error("failed to save cursor", "error", err, "relayUrl", fp.relayUrl) } + return nil // don't exit, just log error + }) + + // save cursor one last time on shutdown + if err := fp.saveCursor(ctx); err != nil { + fp.logger.Error("failed to save cursor on shutdown", "error", err, "relayUrl", fp.relayUrl) } } diff --git a/cmd/tap/util.go b/cmd/tap/util.go index 76a0ff27..98415ca7 100644 --- a/cmd/tap/util.go +++ b/cmd/tap/util.go @@ -1,6 +1,7 @@ package main import ( + "context" "errors" "fmt" "math/rand" @@ -82,3 +83,20 @@ func parseOutboxMode(webhookURL string, disableAcks bool) OutboxMode { return OutboxModeWebsocketAck } } + +// runPeriodically runs the provided task function at the specified interval until the context is done or an error occurs. +func runPeriodically(ctx context.Context, interval time.Duration, task func(context.Context) error) error { + ticker := time.NewTicker(interval) + defer ticker.Stop() + + for { + select { + case <-ctx.Done(): + return ctx.Err() + case <-ticker.C: + if err := task(ctx); err != nil { + return fmt.Errorf("periodic task failed: %w", err) + } + } + } +} diff --git a/cmd/tap/util_test.go b/cmd/tap/util_test.go new file mode 100644 index 00000000..c20a92b4 --- /dev/null +++ b/cmd/tap/util_test.go @@ -0,0 +1,71 @@ +package main + +import ( + "context" + "errors" + "strings" + "sync/atomic" + "testing" + "time" +) + +func TestRunPeriodicallyStopsOnTaskError(t *testing.T) { + + ctx := t.Context() + + var calls atomic.Int32 + taskErr := errors.New("boom") + + errCh := make(chan error, 1) + go func() { + errCh <- runPeriodically(ctx, 5*time.Millisecond, func(context.Context) error { + calls.Add(1) + return taskErr + }) + }() + + select { + case err := <-errCh: + if !errors.Is(err, taskErr) { + t.Fatalf("expected error to wrap task error, got %v", err) + } + if !strings.Contains(err.Error(), "periodic task failed") { + t.Fatalf("expected wrapped error message, got %v", err) + } + if calls.Load() != 1 { + t.Fatalf("expected task to be called once, got %d", calls.Load()) + } + case <-time.After(200 * time.Millisecond): + t.Fatal("runPeriodically did not return on task error") + } +} + +func TestRunPeriodicallyReturnsOnContextCancel(t *testing.T) { + + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + + var calls atomic.Int32 + + errCh := make(chan error, 1) + go func() { + errCh <- runPeriodically(ctx, 5*time.Millisecond, func(context.Context) error { + if calls.Add(1) >= 2 { + cancel() + } + return nil + }) + }() + + select { + case err := <-errCh: + if !errors.Is(err, context.Canceled) { + t.Fatalf("expected context cancellation, got %v", err) + } + if calls.Load() < 2 { + t.Fatalf("expected at least two task executions, got %d", calls.Load()) + } + case <-time.After(250 * time.Millisecond): + t.Fatal("runPeriodically did not return on context cancel") + } +}