Skip to content

Commit c2c89ae

Browse files
committed
Add write retry
1 parent f5bef5d commit c2c89ae

5 files changed

Lines changed: 199 additions & 7 deletions

File tree

‎internal/connectors/kafka.go‎

Lines changed: 11 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -37,6 +37,7 @@ import (
3737
"github.com/IBM/sarama"
3838
v1 "github.com/dataflow-operator/dataflow/api/v1"
3939
"github.com/dataflow-operator/dataflow/internal/metrics"
40+
"github.com/dataflow-operator/dataflow/internal/retry"
4041
"github.com/dataflow-operator/dataflow/internal/types"
4142
"github.com/go-logr/logr"
4243
"github.com/hamba/avro/v2"
@@ -954,7 +955,16 @@ func (k *KafkaSinkConnector) Write(ctx context.Context, messages <-chan *types.M
954955
kafkaMsg.Key = sarama.StringEncoder(key)
955956
}
956957

957-
partition, offset, err := k.producer.SendMessage(kafkaMsg)
958+
var partition int32
959+
var offset int64
960+
err := retry.OnTimeout(ctx, retry.DefaultMaxAttempts, retry.DefaultInitialBackoff, func() error {
961+
p, o, sendErr := k.producer.SendMessage(kafkaMsg)
962+
if sendErr != nil {
963+
return sendErr
964+
}
965+
partition, offset = p, o
966+
return nil
967+
})
958968
if err != nil {
959969
// Record error metric
960970
if k.namespace != "" && k.name != "" {

‎internal/connectors/postgresql.go‎

Lines changed: 10 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,7 @@ import (
2424
"time"
2525

2626
v1 "github.com/dataflow-operator/dataflow/api/v1"
27+
"github.com/dataflow-operator/dataflow/internal/retry"
2728
"github.com/dataflow-operator/dataflow/internal/types"
2829
"github.com/jackc/pgx/v5"
2930
)
@@ -308,14 +309,18 @@ func (p *PostgreSQLSinkConnector) Write(ctx context.Context, messages <-chan *ty
308309
case <-ctx.Done():
309310
if batch.Len() > 0 {
310311
fmt.Printf("DEBUG: Executing final batch on context done, size: %d\n", batch.Len())
311-
return p.executeBatch(ctx, batch)
312+
return retry.OnTimeout(ctx, retry.DefaultMaxAttempts, retry.DefaultInitialBackoff, func() error {
313+
return p.executeBatch(ctx, batch)
314+
})
312315
}
313316
return ctx.Err()
314317
case msg, ok := <-messages:
315318
if !ok {
316319
if batch.Len() > 0 {
317320
fmt.Printf("DEBUG: Executing final batch on channel close, size: %d\n", batch.Len())
318-
return p.executeBatch(ctx, batch)
321+
return retry.OnTimeout(ctx, retry.DefaultMaxAttempts, retry.DefaultInitialBackoff, func() error {
322+
return p.executeBatch(ctx, batch)
323+
})
319324
}
320325
fmt.Printf("DEBUG: Channel closed, no batch to execute\n")
321326
return nil
@@ -432,7 +437,9 @@ func (p *PostgreSQLSinkConnector) Write(ctx context.Context, messages <-chan *ty
432437

433438
if count >= batchSize {
434439
fmt.Printf("DEBUG: Batch size reached, executing batch\n")
435-
if err := p.executeBatch(ctx, batch); err != nil {
440+
if err := retry.OnTimeout(ctx, retry.DefaultMaxAttempts, retry.DefaultInitialBackoff, func() error {
441+
return p.executeBatch(ctx, batch)
442+
}); err != nil {
436443
fmt.Printf("ERROR: Batch execution failed: %v\n", err)
437444
return err
438445
}

‎internal/connectors/trino.go‎

Lines changed: 10 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -27,6 +27,7 @@ import (
2727
"time"
2828

2929
v1 "github.com/dataflow-operator/dataflow/api/v1"
30+
"github.com/dataflow-operator/dataflow/internal/retry"
3031
"github.com/dataflow-operator/dataflow/internal/types"
3132
"github.com/go-logr/logr"
3233
"golang.org/x/oauth2"
@@ -1220,14 +1221,18 @@ func (t *TrinoSinkConnector) Write(ctx context.Context, messages <-chan *types.M
12201221
case <-ctx.Done():
12211222
t.logger.Info("Context cancelled, flushing batch", "batchSize", len(batch))
12221223
if len(batch) > 0 {
1223-
return t.executeBatch(ctx, batch)
1224+
return retry.OnTimeout(ctx, retry.DefaultMaxAttempts, retry.DefaultInitialBackoff, func() error {
1225+
return t.executeBatch(ctx, batch)
1226+
})
12241227
}
12251228
return ctx.Err()
12261229
case msg, ok := <-messages:
12271230
if !ok {
12281231
t.logger.Info("Message channel closed, flushing batch", "batchSize", len(batch), "totalMessages", messageCount)
12291232
if len(batch) > 0 {
1230-
return t.executeBatch(ctx, batch)
1233+
return retry.OnTimeout(ctx, retry.DefaultMaxAttempts, retry.DefaultInitialBackoff, func() error {
1234+
return t.executeBatch(ctx, batch)
1235+
})
12311236
}
12321237
return nil
12331238
}
@@ -1239,7 +1244,9 @@ func (t *TrinoSinkConnector) Write(ctx context.Context, messages <-chan *types.M
12391244

12401245
if len(batch) >= batchSize {
12411246
t.logger.Info("Batch size reached, executing batch", "batchSize", len(batch))
1242-
if err := t.executeBatch(ctx, batch); err != nil {
1247+
if err := retry.OnTimeout(ctx, retry.DefaultMaxAttempts, retry.DefaultInitialBackoff, func() error {
1248+
return t.executeBatch(ctx, batch)
1249+
}); err != nil {
12431250
t.logger.Error(err, "Failed to execute batch", "batchSize", len(batch))
12441251
return err
12451252
}

‎internal/retry/retry.go‎

Lines changed: 56 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,56 @@
1+
package retry
2+
3+
import (
4+
"context"
5+
"errors"
6+
"strings"
7+
"time"
8+
)
9+
10+
// DefaultMaxAttempts is the default number of retry attempts for timeout errors.
11+
const DefaultMaxAttempts = 3
12+
13+
// DefaultInitialBackoff is the initial delay between retries.
14+
const DefaultInitialBackoff = 500 * time.Millisecond
15+
16+
// IsTimeoutError returns true if err is or wraps context.DeadlineExceeded,
17+
// or if the error message indicates a timeout (e.g. from drivers).
18+
func IsTimeoutError(err error) bool {
19+
if err == nil {
20+
return false
21+
}
22+
if errors.Is(err, context.DeadlineExceeded) {
23+
return true
24+
}
25+
msg := strings.ToLower(err.Error())
26+
return strings.Contains(msg, "timeout") ||
27+
strings.Contains(msg, "deadline exceeded") ||
28+
strings.Contains(msg, "i/o timeout")
29+
}
30+
31+
// OnTimeout runs op and retries up to maxAttempts times when op returns a timeout error.
32+
// Backoff doubles after each attempt (initialBackoff, 2*initialBackoff, ...).
33+
// If op returns a non-timeout error, it is returned immediately without retry.
34+
func OnTimeout(ctx context.Context, maxAttempts int, initialBackoff time.Duration, op func() error) error {
35+
var lastErr error
36+
backoff := initialBackoff
37+
for attempt := 0; attempt < maxAttempts; attempt++ {
38+
lastErr = op()
39+
if lastErr == nil {
40+
return nil
41+
}
42+
if !IsTimeoutError(lastErr) {
43+
return lastErr
44+
}
45+
if attempt == maxAttempts-1 {
46+
return lastErr
47+
}
48+
select {
49+
case <-ctx.Done():
50+
return ctx.Err()
51+
case <-time.After(backoff):
52+
backoff *= 2
53+
}
54+
}
55+
return lastErr
56+
}

‎internal/retry/retry_test.go‎

Lines changed: 112 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,112 @@
1+
package retry
2+
3+
import (
4+
"context"
5+
"errors"
6+
"testing"
7+
"time"
8+
)
9+
10+
func TestIsTimeoutError(t *testing.T) {
11+
tests := []struct {
12+
name string
13+
err error
14+
want bool
15+
}{
16+
{"nil", nil, false},
17+
{"DeadlineExceeded", context.DeadlineExceeded, true},
18+
{"wrapped DeadlineExceeded", errors.Join(errors.New("wrap"), context.DeadlineExceeded), true},
19+
{"timeout in message", errors.New("connection timeout"), true},
20+
{"Timeout in message", errors.New("connection Timeout"), true},
21+
{"i/o timeout", errors.New("read tcp: i/o timeout"), true},
22+
{"deadline exceeded in message", errors.New("context deadline exceeded"), true},
23+
{"other error", errors.New("something went wrong"), false},
24+
}
25+
for _, tt := range tests {
26+
t.Run(tt.name, func(t *testing.T) {
27+
if got := IsTimeoutError(tt.err); got != tt.want {
28+
t.Errorf("IsTimeoutError() = %v, want %v", got, tt.want)
29+
}
30+
})
31+
}
32+
}
33+
34+
func TestOnTimeout_SuccessFirstTry(t *testing.T) {
35+
ctx := context.Background()
36+
calls := 0
37+
err := OnTimeout(ctx, 3, 10*time.Millisecond, func() error {
38+
calls++
39+
return nil
40+
})
41+
if err != nil {
42+
t.Errorf("OnTimeout() err = %v, want nil", err)
43+
}
44+
if calls != 1 {
45+
t.Errorf("expected 1 call, got %d", calls)
46+
}
47+
}
48+
49+
func TestOnTimeout_NonTimeoutErrorNoRetry(t *testing.T) {
50+
ctx := context.Background()
51+
wantErr := errors.New("permanent error")
52+
calls := 0
53+
err := OnTimeout(ctx, 3, 10*time.Millisecond, func() error {
54+
calls++
55+
return wantErr
56+
})
57+
if !errors.Is(err, wantErr) {
58+
t.Errorf("OnTimeout() err = %v, want %v", err, wantErr)
59+
}
60+
if calls != 1 {
61+
t.Errorf("expected 1 call (no retry on non-timeout), got %d", calls)
62+
}
63+
}
64+
65+
func TestOnTimeout_RetryThenSuccess(t *testing.T) {
66+
ctx := context.Background()
67+
calls := 0
68+
err := OnTimeout(ctx, 3, 5*time.Millisecond, func() error {
69+
calls++
70+
if calls < 2 {
71+
return context.DeadlineExceeded
72+
}
73+
return nil
74+
})
75+
if err != nil {
76+
t.Errorf("OnTimeout() err = %v, want nil", err)
77+
}
78+
if calls != 2 {
79+
t.Errorf("expected 2 calls (retry then success), got %d", calls)
80+
}
81+
}
82+
83+
func TestOnTimeout_ExhaustRetries(t *testing.T) {
84+
ctx := context.Background()
85+
calls := 0
86+
err := OnTimeout(ctx, 3, 5*time.Millisecond, func() error {
87+
calls++
88+
return context.DeadlineExceeded
89+
})
90+
if err != context.DeadlineExceeded {
91+
t.Errorf("OnTimeout() err = %v, want DeadlineExceeded", err)
92+
}
93+
if calls != 3 {
94+
t.Errorf("expected 3 calls (all retries), got %d", calls)
95+
}
96+
}
97+
98+
func TestOnTimeout_ContextCanceled(t *testing.T) {
99+
ctx, cancel := context.WithCancel(context.Background())
100+
cancel()
101+
calls := 0
102+
err := OnTimeout(ctx, 3, 100*time.Millisecond, func() error {
103+
calls++
104+
return context.DeadlineExceeded
105+
})
106+
if err != context.Canceled {
107+
t.Errorf("OnTimeout() err = %v, want context.Canceled", err)
108+
}
109+
if calls != 1 {
110+
t.Errorf("expected 1 call then context cancel, got %d", calls)
111+
}
112+
}

0 commit comments

Comments
 (0)