Skip to content

Commit 7122a1e

Browse files
committed
test(queue): cover handler context propagation across drivers
1 parent 0b390ca commit 7122a1e

5 files changed

Lines changed: 420 additions & 2 deletions

File tree

driver/redisqueue/worker_redis_impl.go

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,7 @@ type redisWorker struct {
2020
server server
2121
mux *backend.ServeMux
2222
obs queue.Observer
23+
ctxDecorator func(context.Context) context.Context
2324

2425
mu sync.Mutex
2526
started bool
@@ -29,17 +30,31 @@ func newRedisWorker(server server, mux *backend.ServeMux, observer queue.Observe
2930
return &redisWorker{server: server, mux: mux, obs: observer}
3031
}
3132

33+
func (w *redisWorker) SetHandlerContextDecorator(fn func(context.Context) context.Context) {
34+
w.ctxDecorator = fn
35+
}
36+
3237
func (w *redisWorker) Register(jobType string, handler queue.Handler) {
3338
if jobType == "" || handler == nil {
3439
return
3540
}
3641
if w.obs == nil {
3742
w.mux.HandleFunc(jobType, func(ctx context.Context, job *backend.Task) error {
43+
if w.ctxDecorator != nil {
44+
if decorated := w.ctxDecorator(ctx); decorated != nil {
45+
ctx = decorated
46+
}
47+
}
3848
return handler(ctx, queue.NewJob(job.Type()).Payload(job.Payload()))
3949
})
4050
return
4151
}
4252
w.mux.HandleFunc(jobType, func(ctx context.Context, job *backend.Task) error {
53+
if w.ctxDecorator != nil {
54+
if decorated := w.ctxDecorator(ctx); decorated != nil {
55+
ctx = decorated
56+
}
57+
}
4358
attempt, _ := backend.GetRetryCount(ctx)
4459
maxRetry, _ := backend.GetMaxRetry(ctx)
4560
queueName, _ := backend.GetQueueName(ctx)

driver/redisqueue/worker_redis_impl_test.go

Lines changed: 47 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -101,7 +101,7 @@ func TestRedisWorker_ShutdownHonorsContext(t *testing.T) {
101101
func TestRedisWorker_ProcessEventsWithObserver(t *testing.T) {
102102
server := &serverStub{}
103103
var events []queue.Event
104-
observer := queue.ObserverFunc(func(event queue.Event) { events = append(events, event) })
104+
observer := queue.ObserverFunc(func(_ context.Context, event queue.Event) { events = append(events, event) })
105105
w := newRedisWorker(server, backend.NewServeMux(), observer)
106106

107107
w.Register("job:ok", func(context.Context, queue.Job) error { return nil })
@@ -156,7 +156,7 @@ func TestRedisWorker_ProcessEventsWithObserver(t *testing.T) {
156156
func TestRedisWorker_ProcessEventsUnwrapBusEnvelopeJobType(t *testing.T) {
157157
server := &serverStub{}
158158
var events []queue.Event
159-
observer := queue.ObserverFunc(func(event queue.Event) { events = append(events, event) })
159+
observer := queue.ObserverFunc(func(_ context.Context, event queue.Event) { events = append(events, event) })
160160
w := newRedisWorker(server, backend.NewServeMux(), observer)
161161

162162
w.Register("bus:job", func(context.Context, queue.Job) error { return nil })
@@ -210,3 +210,48 @@ func TestRedisWorker_NoObserverFastPath(t *testing.T) {
210210
t.Fatalf("expected handler called once, got %d", called)
211211
}
212212
}
213+
214+
func TestRedisWorker_ObserverSeesDecoratedContext(t *testing.T) {
215+
server := &serverStub{}
216+
type ctxKey struct{}
217+
key := ctxKey{}
218+
const want = "jobs"
219+
220+
var observed []string
221+
var handled []string
222+
observer := queue.ObserverFunc(func(ctx context.Context, event queue.Event) {
223+
if event.Kind != queue.EventProcessStarted && event.Kind != queue.EventProcessSucceeded {
224+
return
225+
}
226+
value, _ := ctx.Value(key).(string)
227+
observed = append(observed, value)
228+
})
229+
w := newRedisWorker(server, backend.NewServeMux(), observer)
230+
w.SetHandlerContextDecorator(func(ctx context.Context) context.Context {
231+
return context.WithValue(ctx, key, want)
232+
})
233+
234+
w.Register("job:decorated", func(ctx context.Context, _ queue.Job) error {
235+
value, _ := ctx.Value(key).(string)
236+
handled = append(handled, value)
237+
return nil
238+
})
239+
if err := w.StartWorkers(context.Background()); err != nil {
240+
t.Fatalf("start workers failed: %v", err)
241+
}
242+
if err := server.lastStartHandler.ProcessTask(context.Background(), backend.NewTask("job:decorated", []byte("ok"))); err != nil {
243+
t.Fatalf("process task failed: %v", err)
244+
}
245+
246+
if len(observed) != 2 {
247+
t.Fatalf("expected 2 observed events, got %d", len(observed))
248+
}
249+
for i, got := range observed {
250+
if got != want {
251+
t.Fatalf("expected observed[%d] = %q, got %q", i, want, got)
252+
}
253+
}
254+
if len(handled) != 1 || handled[0] != want {
255+
t.Fatalf("expected handler to see %q, got %#v", want, handled)
256+
}
257+
}

0 commit comments

Comments
 (0)