diff --git a/v2/inflightmessages.go b/v2/inflightmessages.go new file mode 100644 index 0000000..0cdeaec --- /dev/null +++ b/v2/inflightmessages.go @@ -0,0 +1,79 @@ +package shuttle + +import ( + "context" + "errors" + "fmt" + "sync" + + "github.com/Azure/azure-sdk-for-go/sdk/messaging/azservicebus" +) + +type inFlightMessages struct { + mu sync.RWMutex + tracked map[*azservicebus.ReceivedMessage]struct{} +} + +func newInFlightMessages() *inFlightMessages { + return &inFlightMessages{ + tracked: make(map[*azservicebus.ReceivedMessage]struct{}), + } +} + +func (m *inFlightMessages) track(message *azservicebus.ReceivedMessage) { + m.mu.Lock() + defer m.mu.Unlock() + m.tracked[message] = struct{}{} +} + +func (m *inFlightMessages) forget(message *azservicebus.ReceivedMessage) { + m.mu.Lock() + defer m.mu.Unlock() + delete(m.tracked, message) +} + +func (m *inFlightMessages) messages() []*azservicebus.ReceivedMessage { + m.mu.RLock() + defer m.mu.RUnlock() + + messages := make([]*azservicebus.ReceivedMessage, 0, len(m.tracked)) + for message := range m.tracked { + messages = append(messages, message) + } + return messages +} + +func (m *inFlightMessages) close(ctx context.Context, settler MessageSettler) error { + messages := m.messages() + abandonResults := make(chan error, len(messages)) + + for _, message := range messages { + message := message + go func() { + abandonResults <- m.abandon(ctx, settler, message) + }() + } + + var errs []error + for range messages { + select { + case err := <-abandonResults: + if err != nil { + errs = append(errs, err) + } + case <-ctx.Done(): + getLogger(ctx).Warn(fmt.Sprintf("processor close context done while abandoning in-flight messages: %s", ctx.Err())) + return errors.Join(append(errs, ctx.Err())...) + } + } + return errors.Join(errs...) +} + +func (m *inFlightMessages) abandon(ctx context.Context, settler MessageSettler, message *azservicebus.ReceivedMessage) error { + err := settler.AbandonMessage(ctx, message, nil) + if err != nil { + return fmt.Errorf("failed to abandon message %s during processor close: %w", message.MessageID, err) + } + m.forget(message) + return nil +} diff --git a/v2/inflightmessages_test.go b/v2/inflightmessages_test.go new file mode 100644 index 0000000..311d29e --- /dev/null +++ b/v2/inflightmessages_test.go @@ -0,0 +1,262 @@ +package shuttle + +import ( + "context" + "errors" + "sync" + "testing" + "time" + + "github.com/Azure/azure-sdk-for-go/sdk/messaging/azservicebus" + "github.com/stretchr/testify/require" +) + +func TestInFlightMessages_CloseAbandonsTrackedMessages(t *testing.T) { + inFlight := newInFlightMessages() + settler := &inFlightMessageSettler{} + trackedMessage := &azservicebus.ReceivedMessage{MessageID: "tracked"} + forgottenMessage := &azservicebus.ReceivedMessage{MessageID: "forgotten"} + + inFlight.track(trackedMessage) + inFlight.track(forgottenMessage) + inFlight.forget(forgottenMessage) + + require.NoError(t, inFlight.close(context.Background(), settler)) + require.Empty(t, inFlight.messages()) + require.NoError(t, inFlight.close(context.Background(), settler)) + + messages := settler.abandonedMessages() + require.Len(t, messages, 1) + require.Same(t, trackedMessage, messages[0]) +} + +func TestInFlightMessages_CloseAbandonsTrackedMessagesConcurrently(t *testing.T) { + inFlight := newInFlightMessages() + firstMessage := &azservicebus.ReceivedMessage{MessageID: "first"} + secondMessage := &azservicebus.ReceivedMessage{MessageID: "second"} + abandonStarted := make(chan *azservicebus.ReceivedMessage, 2) + releaseAbandons := make(chan struct{}) + var releaseOnce sync.Once + settler := &inFlightMessageSettler{ + abandon: func(ctx context.Context, message *azservicebus.ReceivedMessage) error { + abandonStarted <- message + <-releaseAbandons + return nil + }, + } + + release := func() { + releaseOnce.Do(func() { + close(releaseAbandons) + }) + } + t.Cleanup(release) + + inFlight.track(firstMessage) + inFlight.track(secondMessage) + + closeErr := make(chan error, 1) + go func() { + closeErr <- inFlight.close(context.Background(), settler) + }() + + abandoned := map[*azservicebus.ReceivedMessage]struct{}{} + for len(abandoned) < 2 { + select { + case message := <-abandonStarted: + abandoned[message] = struct{}{} + case <-time.After(1 * time.Second): + t.Fatalf("timed out waiting for concurrent abandon attempts") + } + } + + select { + case err := <-closeErr: + t.Fatalf("close returned before abandon attempts were released: %v", err) + default: + } + + release() + require.NoError(t, <-closeErr) + require.Contains(t, abandoned, firstMessage) + require.Contains(t, abandoned, secondMessage) +} + +func TestInFlightMessages_CloseReturnsWhenContextIsCanceled(t *testing.T) { + inFlight := newInFlightMessages() + abandonStarted := make(chan *azservicebus.ReceivedMessage, 1) + releaseAbandon := make(chan struct{}) + var releaseOnce sync.Once + settler := &inFlightMessageSettler{ + abandon: func(ctx context.Context, message *azservicebus.ReceivedMessage) error { + abandonStarted <- message + <-releaseAbandon + return nil + }, + } + + release := func() { + releaseOnce.Do(func() { + close(releaseAbandon) + }) + } + t.Cleanup(release) + + inFlight.track(&azservicebus.ReceivedMessage{MessageID: "blocked"}) + + ctx, cancel := context.WithCancel(context.Background()) + closeErr := make(chan error, 1) + go func() { + closeErr <- inFlight.close(ctx, settler) + }() + + select { + case <-abandonStarted: + case <-time.After(1 * time.Second): + t.Fatal("timed out waiting for abandon attempt to start") + } + + cancel() + + select { + case err := <-closeErr: + require.ErrorIs(t, err, context.Canceled) + case <-time.After(1 * time.Second): + t.Fatal("timed out waiting for close to return after context cancellation") + } + + abandonedMessages := settler.abandonedMessages() + require.Len(t, abandonedMessages, 1) + require.Equal(t, "blocked", abandonedMessages[0].MessageID) + + release() + require.Eventually(t, func() bool { + return len(inFlight.messages()) == 0 + }, 1*time.Second, 10*time.Millisecond) +} + +func TestInFlightMessages_CloseReturnsWhenContextDeadlineExpires(t *testing.T) { + inFlight := newInFlightMessages() + abandonStarted := make(chan *azservicebus.ReceivedMessage, 1) + releaseAbandon := make(chan struct{}) + var releaseOnce sync.Once + settler := &inFlightMessageSettler{ + abandon: func(ctx context.Context, message *azservicebus.ReceivedMessage) error { + abandonStarted <- message + <-releaseAbandon + return nil + }, + } + + release := func() { + releaseOnce.Do(func() { + close(releaseAbandon) + }) + } + t.Cleanup(release) + + inFlight.track(&azservicebus.ReceivedMessage{MessageID: "deadline"}) + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Millisecond) + defer cancel() + + closeErr := make(chan error, 1) + go func() { + closeErr <- inFlight.close(ctx, settler) + }() + + select { + case <-abandonStarted: + case <-time.After(1 * time.Second): + t.Fatal("timed out waiting for abandon attempt to start") + } + + select { + case err := <-closeErr: + require.ErrorIs(t, err, context.DeadlineExceeded) + case <-time.After(1 * time.Second): + t.Fatal("timed out waiting for close to return after context deadline") + } + + abandonedMessages := settler.abandonedMessages() + require.Len(t, abandonedMessages, 1) + require.Equal(t, "deadline", abandonedMessages[0].MessageID) + + release() + require.Eventually(t, func() bool { + return len(inFlight.messages()) == 0 + }, 1*time.Second, 10*time.Millisecond) +} + +func TestInFlightMessages_CloseReturnsAbandonErrorsAndKeepsMessages(t *testing.T) { + firstErr := errors.New("first abandon failed") + secondErr := errors.New("second abandon failed") + inFlight := newInFlightMessages() + firstMessage := &azservicebus.ReceivedMessage{MessageID: "first"} + secondMessage := &azservicebus.ReceivedMessage{MessageID: "second"} + abandonErrors := make(chan error, 2) + abandonErrors <- firstErr + abandonErrors <- secondErr + settler := &inFlightMessageSettler{ + abandon: func(ctx context.Context, message *azservicebus.ReceivedMessage) error { + return <-abandonErrors + }, + } + + inFlight.track(firstMessage) + inFlight.track(secondMessage) + + err := inFlight.close(context.Background(), settler) + + require.Error(t, err) + require.ErrorIs(t, err, firstErr) + require.ErrorIs(t, err, secondErr) + require.ElementsMatch(t, []*azservicebus.ReceivedMessage{firstMessage, secondMessage}, inFlight.messages()) + + successfulSettler := &inFlightMessageSettler{} + require.NoError(t, inFlight.close(context.Background(), successfulSettler)) + require.Empty(t, inFlight.messages()) + abandonedMessages := successfulSettler.abandonedMessages() + require.ElementsMatch(t, []*azservicebus.ReceivedMessage{firstMessage, secondMessage}, abandonedMessages) +} + +type inFlightMessageSettler struct { + mu sync.Mutex + abandoned []*azservicebus.ReceivedMessage + abandon func(context.Context, *azservicebus.ReceivedMessage) error +} + +func (s *inFlightMessageSettler) AbandonMessage(ctx context.Context, message *azservicebus.ReceivedMessage, options *azservicebus.AbandonMessageOptions) error { + s.mu.Lock() + s.abandoned = append(s.abandoned, message) + s.mu.Unlock() + + if s.abandon != nil { + return s.abandon(ctx, message) + } + return nil +} + +func (s *inFlightMessageSettler) CompleteMessage(ctx context.Context, message *azservicebus.ReceivedMessage, options *azservicebus.CompleteMessageOptions) error { + return nil +} + +func (s *inFlightMessageSettler) DeadLetterMessage(ctx context.Context, message *azservicebus.ReceivedMessage, options *azservicebus.DeadLetterOptions) error { + return nil +} + +func (s *inFlightMessageSettler) DeferMessage(ctx context.Context, message *azservicebus.ReceivedMessage, options *azservicebus.DeferMessageOptions) error { + return nil +} + +func (s *inFlightMessageSettler) RenewMessageLock(ctx context.Context, message *azservicebus.ReceivedMessage, options *azservicebus.RenewMessageLockOptions) error { + return nil +} + +func (s *inFlightMessageSettler) abandonedMessages() []*azservicebus.ReceivedMessage { + s.mu.Lock() + defer s.mu.Unlock() + messages := make([]*azservicebus.ReceivedMessage, len(s.abandoned)) + copy(messages, s.abandoned) + return messages +} diff --git a/v2/processor.go b/v2/processor.go index 814de7b..79aa149 100644 --- a/v2/processor.go +++ b/v2/processor.go @@ -4,6 +4,7 @@ import ( "context" "errors" "fmt" + "sync" "time" "github.com/Azure/azure-sdk-for-go/sdk/azcore/to" @@ -43,6 +44,50 @@ type Processor struct { options ProcessorOptions handle Handler concurrencyTokens chan struct{} // tracks how many concurrent messages are currently being handled by the processor + inFlightMessages *inFlightMessages + lifecycle processorLifecycle +} + +type processorLifecycle struct { + mu sync.Mutex + closed bool + cancelStartCtx context.CancelFunc +} + +func (l *processorLifecycle) start(ctx context.Context) (context.Context, context.CancelFunc) { + l.mu.Lock() + defer l.mu.Unlock() + + if l.closed { + return nil, nil + } + + ctx, cancelStartCtx := context.WithCancel(ctx) + l.cancelStartCtx = cancelStartCtx + return ctx, cancelStartCtx +} + +func (l *processorLifecycle) close() { + if cancelStartCtx := l.markClosed(); cancelStartCtx != nil { + cancelStartCtx() + } +} + +func (l *processorLifecycle) markClosed() context.CancelFunc { + l.mu.Lock() + defer l.mu.Unlock() + + l.closed = true + cancelStartCtx := l.cancelStartCtx + l.cancelStartCtx = nil + return cancelStartCtx +} + +func (l *processorLifecycle) clearCancelStartCtx() { + l.mu.Lock() + defer l.mu.Unlock() + + l.cancelStartCtx = nil } // ProcessorOptions configures the processor @@ -106,6 +151,7 @@ func NewProcessor(receiver Receiver, handler HandlerFunc, options *ProcessorOpti handle: handler, options: *opts, concurrencyTokens: make(chan struct{}, opts.MaxConcurrency), + inFlightMessages: newInFlightMessages(), } } @@ -113,6 +159,12 @@ func NewProcessor(receiver Receiver, handler HandlerFunc, options *ProcessorOpti // It will retry starting the processor based on the StartMaxAttempt and StartRetryDelayStrategy. // Returns a combined list of errors encountered during each processor start attempt. func (p *Processor) Start(ctx context.Context) (err error) { + ctx, cancel := p.lifecycle.start(ctx) + if ctx == nil { + return fmt.Errorf("failed to start processor since it has already been closed: %w", context.Canceled) + } + defer p.lifecycle.clearCancelStartCtx() + defer cancel() defer func() { if rec := recover(); rec != nil { err = fmt.Errorf("panic recovered from processor: %s", rec) @@ -121,6 +173,15 @@ func (p *Processor) Start(ctx context.Context) (err error) { return p.startWithRetries(ctx) } +// Close stops receiving new messages, cancels in-flight message handlers, and +// best-effort abandons messages currently held by the processor. +func (p *Processor) Close(ctx context.Context) error { + // Close the active processor loop so no further messages are received. + p.lifecycle.close() + + return p.inFlightMessages.close(ctx, p.receiver) +} + // startWithRetries starts a processor and blocks until an error occurs or the context is canceled. // It will retry starting the processor based on the StartMaxAttempt and StartRetryDelayStrategy. // Returns a combined list of errors during the start attempts or ctx.Err() if the context @@ -151,14 +212,8 @@ func (p *Processor) startWithRetries(ctx context.Context) error { func (p *Processor) start(ctx context.Context) error { logger := getLogger(ctx) logger.Info("starting processor") - messages, err := p.receiver.ReceiveMessages(ctx, p.options.MaxReceiveCount, nil) - if err != nil { - return fmt.Errorf("failed to receive messages: %w", err) - } - logger.Info(fmt.Sprintf("received %d messages - initial", len(messages))) - processor.Metric.IncMessageReceived(float64(len(messages))) - for _, msg := range messages { - p.process(ctx, msg) + if err := p.receiveAndProcess(ctx, p.options.MaxReceiveCount, "initial"); err != nil { + return err } for ctx.Err() == nil { select { @@ -167,14 +222,8 @@ func (p *Processor) start(ctx context.Context) error { if ctx.Err() != nil || maxMessages == 0 { break } - messages, err := p.receiver.ReceiveMessages(ctx, maxMessages, nil) - if err != nil { - return fmt.Errorf("failed to receive messages: %w", err) - } - logger.Info(fmt.Sprintf("received %d messages from processor loop", len(messages))) - processor.Metric.IncMessageReceived(float64(len(messages))) - for _, msg := range messages { - p.process(ctx, msg) + if err := p.receiveAndProcess(ctx, maxMessages, "processor loop"); err != nil { + return err } case <-ctx.Done(): logger.Info("context done, stop receiving from processor") @@ -184,13 +233,29 @@ func (p *Processor) start(ctx context.Context) error { return ctx.Err() } +func (p *Processor) receiveAndProcess(ctx context.Context, maxMessages int, source string) error { + messages, err := p.receiver.ReceiveMessages(ctx, maxMessages, nil) + if err != nil { + return fmt.Errorf("failed to receive messages: %w", err) + } + getLogger(ctx).Info(fmt.Sprintf("received %d messages - %s", len(messages), source)) + processor.Metric.IncMessageReceived(float64(len(messages))) + for _, msg := range messages { + p.process(ctx, msg) + } + return nil +} + func (p *Processor) process(ctx context.Context, message *azservicebus.ReceivedMessage) { p.concurrencyTokens <- struct{}{} + p.inFlightMessages.track(message) + go func() { msgContext, cancel := context.WithCancel(ctx) // cancel messageContext when we get out of this goroutine defer cancel() defer func() { + p.inFlightMessages.forget(message) <-p.concurrencyTokens processor.Metric.IncMessageHandled(message) processor.Metric.DecConcurrentMessageCount(message) diff --git a/v2/processor_fake_test.go b/v2/processor_fake_test.go index 4c90798..c0d340b 100644 --- a/v2/processor_fake_test.go +++ b/v2/processor_fake_test.go @@ -18,11 +18,12 @@ type fakeSettler struct { DeadLetterCalled atomic.Int32 DeferCalled atomic.Int32 RenewCalled atomic.Int32 + SetupAbandonErr error } func (f *fakeSettler) AbandonMessage(ctx context.Context, message *azservicebus.ReceivedMessage, options *azservicebus.AbandonMessageOptions) error { f.AbandonCalled.Add(1) - return nil + return f.SetupAbandonErr } func (f *fakeSettler) CompleteMessage(ctx context.Context, message *azservicebus.ReceivedMessage, options *azservicebus.CompleteMessageOptions) error { @@ -55,21 +56,40 @@ type fakeReceiver struct { *fakeSettler SetupMaxReceiveCalls int SetupReceivePanic string + SetupReceiveStarted chan struct{} } -func (f *fakeReceiver) ReceiveMessages(_ context.Context, maxMessages int, _ *azservicebus.ReceiveMessagesOptions) ([]*azservicebus.ReceivedMessage, error) { +func (f *fakeReceiver) ReceiveMessages(ctx context.Context, maxMessages int, _ *azservicebus.ReceiveMessagesOptions) ([]*azservicebus.ReceivedMessage, error) { f.ReceiveCalls = append(f.ReceiveCalls, maxMessages) + if f.SetupReceiveStarted != nil { + select { + case f.SetupReceiveStarted <- struct{}{}: + default: + } + } if maxMessages == 0 && len(f.SetupReceivedMessages) > 0 { return nil, nil } var result []*azservicebus.ReceivedMessage - for msg := range f.SetupReceivedMessages { - result = append(result, msg) - if len(result) == maxMessages || len(f.SetupReceivedMessages) == 0 { - break + for len(result) < maxMessages { + select { + case msg, ok := <-f.SetupReceivedMessages: + if !ok { + return f.receiveResult(result) + } + result = append(result, msg) + if len(f.SetupReceivedMessages) == 0 { + return f.receiveResult(result) + } + case <-ctx.Done(): + return result, ctx.Err() } } + return f.receiveResult(result) +} + +func (f *fakeReceiver) receiveResult(result []*azservicebus.ReceivedMessage) ([]*azservicebus.ReceivedMessage, error) { if f.SetupReceivePanic != "" { panic(f.SetupReceivePanic) } diff --git a/v2/processor_test.go b/v2/processor_test.go index 9e2c3d3..7ae2932 100644 --- a/v2/processor_test.go +++ b/v2/processor_test.go @@ -2,6 +2,7 @@ package shuttle_test import ( "context" + "errors" "fmt" "testing" "time" @@ -69,9 +70,7 @@ func TestProcessorStart_DefaultsToMaxConcurrency(t *testing.T) { SetupReceivedMessages: messages, } processor := shuttle.NewProcessor(rcv, MyHandler(0*time.Second), nil) - ctx, cancel := context.WithCancel(context.Background()) - cancel() - err := processor.Start(ctx) + err := processor.Start(context.Background()) a.ErrorContains(err, "max receive calls exceeded") a.Equal(1, len(rcv.ReceiveCalls), "there should be 1 entry in the ReceiveCalls array") a.Equal(1, rcv.ReceiveCalls[0], "the processor should have used the default max concurrency of 1") @@ -100,6 +99,22 @@ func TestProcessorStart_ContextCanceledAfterStart(t *testing.T) { g.Eventually(errCh).Should(Receive(MatchError(context.Canceled))) } +func TestProcessorStart_ReturnsProcessorClosedErrorAfterClose(t *testing.T) { + rcv := &fakeReceiver{ + fakeSettler: &fakeSettler{}, + SetupReceivedMessages: make(chan *azservicebus.ReceivedMessage), + } + close(rcv.SetupReceivedMessages) + processor := shuttle.NewProcessor(rcv, MyHandler(0*time.Millisecond), nil) + + require.NoError(t, processor.Close(context.Background())) + + err := processor.Start(context.Background()) + + require.ErrorContains(t, err, "failed to start processor since it has already been closed") + require.Empty(t, rcv.ReceiveCalls) +} + func TestProcessorStart_CanSetMaxConcurrency(t *testing.T) { a := require.New(t) rcv := &fakeReceiver{ @@ -110,10 +125,7 @@ func TestProcessorStart_CanSetMaxConcurrency(t *testing.T) { processor := shuttle.NewProcessor(rcv, MyHandler(0*time.Second), &shuttle.ProcessorOptions{ MaxConcurrency: 10, }) - ctx, cancel := context.WithCancel(context.Background()) - // pre-cancel the context - cancel() - err := processor.Start(ctx) + err := processor.Start(context.Background()) a.ErrorContains(err, "max receive calls exceeded") a.Equal(1, len(rcv.ReceiveCalls), "there should be 1 entry in the ReceiveCalls array") a.Equal(10, rcv.ReceiveCalls[0], "the processor should have used max concurrency of 10") @@ -174,9 +186,7 @@ func TestProcessorStart_MaxReceiveCountGreaterThanMaxConcurrency(t *testing.T) { MaxConcurrency: 5, MaxReceiveCount: 10, }) - ctx, cancel := context.WithCancel(context.Background()) - cancel() - err := processor.Start(ctx) + err := processor.Start(context.Background()) a.ErrorContains(err, "max receive calls exceeded") a.Equal(1, len(rcv.ReceiveCalls), "there should be 1 entry in the ReceiveCalls array") a.Equal(5, rcv.ReceiveCalls[0], "the processor should have used max concurrency as max receive count") @@ -193,9 +203,7 @@ func TestProcessorStart_DisableMaxReceiveCount(t *testing.T) { MaxConcurrency: 5, MaxReceiveCount: -1, }) - ctx, cancel := context.WithCancel(context.Background()) - cancel() - err := processor.Start(ctx) + err := processor.Start(context.Background()) a.ErrorContains(err, "max receive calls exceeded") a.Equal(1, len(rcv.ReceiveCalls), "there should be 1 entry in the ReceiveCalls array") a.Equal(5, rcv.ReceiveCalls[0], "the processor should have used max concurrency as max receive count") @@ -294,9 +302,7 @@ func TestProcessorStart_DefaultsToStartMaxAttempt(t *testing.T) { SetupReceiveError: fmt.Errorf("fake receive error"), } processor := shuttle.NewProcessor(rcv, MyHandler(0*time.Second), nil) - ctx, cancel := context.WithCancel(context.Background()) - cancel() - err := processor.Start(ctx) + err := processor.Start(context.Background()) a.ErrorContains(err, "fake receive error") a.Equal(1, len(rcv.ReceiveCalls), "there should be 1 entry in the ReceiveCalls array") a.Equal(1, rcv.ReceiveCalls[0], "the processor should have used the default max concurrency of 1") @@ -360,6 +366,28 @@ func TestProcessorStart_ContextCanceledDuringStartRetry(t *testing.T) { a.Equal(1, rcv.ReceiveCalls[1], "the processor should have retried the receive call once") } +func TestProcessorStart_DoesNotRetryAfterCloseDuringFirstAttempt(t *testing.T) { + rcv := &fakeReceiver{ + fakeSettler: &fakeSettler{}, + SetupReceivedMessages: make(chan *azservicebus.ReceivedMessage), + SetupMaxReceiveCalls: 10, + SetupReceiveStarted: make(chan struct{}, 1), + } + + processor := shuttle.NewProcessor(rcv, MyHandler(0*time.Second), &shuttle.ProcessorOptions{ + StartMaxAttempt: 3, + StartRetryDelayStrategy: &shuttle.ConstantDelayStrategy{Delay: 20 * time.Millisecond}, + }) + errCh := make(chan error, 1) + go func() { errCh <- processor.Start(context.Background()) }() + + g := NewWithT(t) + g.Eventually(rcv.SetupReceiveStarted).Should(Receive()) + g.Expect(processor.Close(context.Background())).To(Succeed()) + g.Eventually(errCh).Should(Receive(MatchError(context.Canceled))) + g.Expect(rcv.ReceiveCalls).To(HaveLen(1)) +} + func TestProcessorStart_RecoversReceiverPanic(t *testing.T) { rcv := &fakeReceiver{ fakeSettler: &fakeSettler{}, @@ -377,6 +405,156 @@ func TestProcessorStart_RecoversReceiverPanic(t *testing.T) { g.Expect(err.Error()).To(ContainSubstring("panic recovered from processor: receive panic!")) } +func TestProcessorClose_CancelsAndAbandonsInflightMessages(t *testing.T) { + messages := messagesChannel(2) + close(messages) + settler := &fakeSettler{} + rcv := &fakeReceiver{ + fakeSettler: settler, + SetupReceivedMessages: messages, + SetupMaxReceiveCalls: 10, + } + started := make(chan *azservicebus.ReceivedMessage, 2) + canceled := make(chan *azservicebus.ReceivedMessage, 2) + processor := shuttle.NewProcessor(rcv, func(ctx context.Context, settler shuttle.MessageSettler, message *azservicebus.ReceivedMessage) { + started <- message + <-ctx.Done() + canceled <- message + }, &shuttle.ProcessorOptions{ + MaxConcurrency: 2, + ReceiveInterval: to.Ptr(1 * time.Hour), + }) + + errCh := make(chan error, 1) + go func() { errCh <- processor.Start(context.Background()) }() + + g := NewWithT(t) + g.Eventually(started).Should(Receive()) + g.Eventually(started).Should(Receive()) + + g.Expect(processor.Close(context.Background())).To(Succeed()) + g.Eventually(canceled).Should(Receive()) + g.Eventually(canceled).Should(Receive()) + g.Eventually(errCh).Should(Receive(MatchError(context.Canceled))) + g.Expect(settler.AbandonCalled.Load()).To(Equal(int32(2))) +} + +func TestProcessorClose_StopsStartReceiveLoop(t *testing.T) { + rcv := &fakeReceiver{ + fakeSettler: &fakeSettler{}, + SetupReceivedMessages: make(chan *azservicebus.ReceivedMessage), + SetupMaxReceiveCalls: 10, + SetupReceiveStarted: make(chan struct{}, 1), + } + processor := shuttle.NewProcessor(rcv, MyHandler(0*time.Second), &shuttle.ProcessorOptions{ + ReceiveInterval: to.Ptr(10 * time.Millisecond), + }) + errCh := make(chan error, 1) + go func() { errCh <- processor.Start(context.Background()) }() + + g := NewWithT(t) + g.Eventually(rcv.SetupReceiveStarted).Should(Receive()) + g.Expect(processor.Close(context.Background())).To(Succeed()) + g.Eventually(errCh).Should(Receive(MatchError(MatchRegexp("failed to receive messages: context canceled")))) + g.Expect(rcv.ReceiveCalls).To(HaveLen(1)) +} + +func TestProcessorClose_ReturnsAbandonErrors(t *testing.T) { + abandonErr := errors.New("abandon failed") + messages := messagesChannel(2) + close(messages) + settler := &fakeSettler{SetupAbandonErr: abandonErr} + rcv := &fakeReceiver{ + fakeSettler: settler, + SetupReceivedMessages: messages, + SetupMaxReceiveCalls: 10, + } + started := make(chan struct{}, 2) + release := make(chan struct{}) + processor := shuttle.NewProcessor(rcv, func(ctx context.Context, settler shuttle.MessageSettler, message *azservicebus.ReceivedMessage) { + started <- struct{}{} + <-release + }, &shuttle.ProcessorOptions{ + MaxConcurrency: 2, + ReceiveInterval: to.Ptr(1 * time.Hour), + }) + errCh := make(chan error, 1) + go func() { errCh <- processor.Start(context.Background()) }() + defer func() { + close(release) + <-errCh + }() + + g := NewWithT(t) + g.Eventually(started).Should(Receive()) + g.Eventually(started).Should(Receive()) + + err := processor.Close(context.Background()) + g.Expect(err).To(HaveOccurred()) + g.Expect(errors.Is(err, abandonErr)).To(BeTrue()) + g.Expect(err.Error()).To(ContainSubstring("failed to abandon message")) + g.Expect(settler.AbandonCalled.Load()).To(Equal(int32(2))) +} + +func TestProcessorClose_IsIdempotent(t *testing.T) { + messages := messagesChannel(1) + close(messages) + settler := &fakeSettler{} + rcv := &fakeReceiver{ + fakeSettler: settler, + SetupReceivedMessages: messages, + SetupMaxReceiveCalls: 10, + } + started := make(chan struct{}, 1) + processor := shuttle.NewProcessor(rcv, func(ctx context.Context, settler shuttle.MessageSettler, message *azservicebus.ReceivedMessage) { + started <- struct{}{} + <-ctx.Done() + }, &shuttle.ProcessorOptions{ + ReceiveInterval: to.Ptr(1 * time.Hour), + }) + errCh := make(chan error, 1) + go func() { errCh <- processor.Start(context.Background()) }() + + g := NewWithT(t) + g.Eventually(started).Should(Receive()) + g.Expect(processor.Close(context.Background())).To(Succeed()) + g.Expect(processor.Close(context.Background())).To(Succeed()) + g.Eventually(errCh).Should(Receive(MatchError(context.Canceled))) + g.Expect(settler.AbandonCalled.Load()).To(Equal(int32(1))) +} + +func TestProcessorClose_DoesNotWaitForHandlerToExit(t *testing.T) { + messages := messagesChannel(1) + close(messages) + settler := &fakeSettler{} + rcv := &fakeReceiver{ + fakeSettler: settler, + SetupReceivedMessages: messages, + SetupMaxReceiveCalls: 10, + } + started := make(chan struct{}, 1) + release := make(chan struct{}) + processor := shuttle.NewProcessor(rcv, func(ctx context.Context, settler shuttle.MessageSettler, message *azservicebus.ReceivedMessage) { + started <- struct{}{} + <-release + }, &shuttle.ProcessorOptions{ + ReceiveInterval: to.Ptr(1 * time.Hour), + }) + errCh := make(chan error, 1) + go func() { errCh <- processor.Start(context.Background()) }() + + g := NewWithT(t) + g.Eventually(started).Should(Receive()) + + closeCtx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond) + defer cancel() + g.Expect(processor.Close(closeCtx)).To(Succeed()) + g.Expect(settler.AbandonCalled.Load()).To(Equal(int32(1))) + g.Eventually(errCh).Should(Receive(MatchError(context.Canceled))) + + close(release) +} + func messagesChannel(messageCount int) chan *azservicebus.ReceivedMessage { messages := make(chan *azservicebus.ReceivedMessage, messageCount) for i := 0; i < messageCount; i++ {