Skip to content

Commit 2f88288

Browse files
akenraolegz
authored andcommitted
GH-3131 Invoke full ChannelInterceptor contract in DefaultPollableMessageSource
Signed-off-by: akenra <37288280+akenra@users.noreply.github.com> Resolves #3253
1 parent 8cb536c commit 2f88288

2 files changed

Lines changed: 279 additions & 13 deletions

File tree

core/spring-cloud-stream/src/main/java/org/springframework/cloud/stream/binder/DefaultPollableMessageSource.java

Lines changed: 58 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -108,35 +108,70 @@ public DefaultPollableMessageSource(
108108

109109
public void setSource(MessageSource<?> source) {
110110
ProxyFactory pf = new ProxyFactory(source);
111-
class ReceiveAdvice implements MethodInterceptor {
112111

113-
private final List<ChannelInterceptor> interceptors = new ArrayList<>();
112+
class ReceiveAdvice implements MethodInterceptor {
114113

115114
@Override
116115
public Object invoke(MethodInvocation invocation) throws Throwable {
117-
Object result = invocation.proceed();
118-
if (result instanceof Message<?> received) {
119-
for (ChannelInterceptor interceptor : this.interceptors) {
120-
received = interceptor.preSend(received, DUMMY_CHANNEL);
121-
if (received == null) {
116+
Message<?> result = null;
117+
Exception completionException = null;
118+
boolean preReceiveCompleted = false;
119+
try {
120+
for (ChannelInterceptor interceptor : interceptors()) {
121+
if (!interceptor.preReceive(DUMMY_CHANNEL)) {
122122
return null;
123123
}
124124
}
125-
return received;
125+
preReceiveCompleted = true;
126+
Object received = invocation.proceed();
127+
if (received instanceof Message<?> message) {
128+
result = message;
129+
for (ChannelInterceptor interceptor : interceptors()) {
130+
result = interceptor.postReceive(result, DUMMY_CHANNEL);
131+
if (result == null) {
132+
return null;
133+
}
134+
}
135+
for (ChannelInterceptor interceptor : interceptors()) {
136+
result = interceptor.preSend(result, DUMMY_CHANNEL);
137+
if (result == null) {
138+
return null;
139+
}
140+
}
141+
}
142+
else {
143+
result = null;
144+
}
145+
return result;
146+
}
147+
catch (Throwable ex) {
148+
completionException = ex instanceof Exception exception ? exception
149+
: new IllegalStateException(ex);
150+
throw ex;
151+
}
152+
finally {
153+
if (preReceiveCompleted) {
154+
for (ChannelInterceptor interceptor : interceptors()) {
155+
interceptor.afterReceiveCompletion(result, DUMMY_CHANNEL,
156+
completionException);
157+
}
158+
}
126159
}
127-
return result;
128160
}
129161

130162
}
131-
final ReceiveAdvice advice = new ReceiveAdvice();
132-
advice.interceptors.addAll(this.interceptors);
163+
133164
NameMatchMethodPointcutAdvisor sourceAdvisor = new NameMatchMethodPointcutAdvisor(
134-
advice);
165+
new ReceiveAdvice());
135166
sourceAdvisor.addMethodName("receive");
136167
pf.addAdvisor(sourceAdvisor);
137168
this.source = (MessageSource<?>) pf.getProxy();
138169
}
139170

171+
private List<ChannelInterceptor> interceptors() {
172+
return List.copyOf(this.interceptors);
173+
}
174+
140175
public void setRetryTemplate(RetryTemplate retryTemplate) {
141176
this.retryTemplate = retryTemplate;
142177
}
@@ -211,6 +246,7 @@ public boolean poll(MessageHandler handler, ParameterizedTypeReference<?> type)
211246
ackCallback = status -> log.warn("No AcknowledgementCallback defined. Status: " + status.name() + " " + message);
212247
}
213248

249+
Exception sendFailure = null;
214250
try {
215251
setAttributesIfNecessary(message);
216252
if (this.retryTemplate == null) {
@@ -236,13 +272,17 @@ public boolean poll(MessageHandler handler, ParameterizedTypeReference<?> type)
236272
}
237273
}
238274
}
275+
for (ChannelInterceptor interceptor : interceptors()) {
276+
interceptor.postSend(message, DUMMY_CHANNEL, true);
277+
}
239278
return true;
240279
}
241280
catch (MessagingException e) {
281+
sendFailure = e;
242282
if (this.retryTemplate == null && !shouldRequeue(e)) {
243283
try {
244284
this.messagingTemplate.send(this.errorChannel,
245-
this.errorMessageStrategy.buildErrorMessage(e, ATTRIBUTES_HOLDER.get()));
285+
this.errorMessageStrategy.buildErrorMessage(e, ATTRIBUTES_HOLDER.get()));
246286
}
247287
catch (MessagingException e1) {
248288
requeueOrNack(message, ackCallback, e1);
@@ -255,6 +295,7 @@ public boolean poll(MessageHandler handler, ParameterizedTypeReference<?> type)
255295
}
256296
}
257297
catch (Exception e) {
298+
sendFailure = e;
258299
AckUtils.autoNack(ackCallback);
259300
if (e instanceof MessageHandlingException messageHandlingException &&
260301
messageHandlingException.getFailedMessage().equals(message)) {
@@ -265,6 +306,10 @@ public boolean poll(MessageHandler handler, ParameterizedTypeReference<?> type)
265306
finally {
266307
ATTRIBUTES_HOLDER.remove();
267308
AckUtils.autoAck(ackCallback);
309+
for (ChannelInterceptor interceptor : interceptors()) {
310+
interceptor.afterSendCompletion(message, DUMMY_CHANNEL,
311+
sendFailure == null, sendFailure);
312+
}
268313
}
269314
}
270315

Lines changed: 221 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,221 @@
1+
/*
2+
* Copyright 2026-present the original author or authors.
3+
*
4+
* Licensed under the Apache License, Version 2.0 (the "License");
5+
* you may not use this file except in compliance with the License.
6+
* You may obtain a copy of the License at
7+
*
8+
* https://www.apache.org/licenses/LICENSE-2.0
9+
*
10+
* Unless required by applicable law or agreed to in writing, software
11+
* distributed under the License is distributed on an "AS IS" BASIS,
12+
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13+
* See the License for the specific language governing permissions and
14+
* limitations under the License.
15+
*/
16+
17+
package org.springframework.cloud.stream.binder;
18+
19+
import java.util.ArrayList;
20+
import java.util.List;
21+
import java.util.concurrent.atomic.AtomicInteger;
22+
23+
import org.junit.jupiter.api.Test;
24+
25+
import org.springframework.integration.channel.DirectChannel;
26+
import org.springframework.integration.core.MessageSource;
27+
import org.springframework.messaging.Message;
28+
import org.springframework.messaging.MessageChannel;
29+
import org.springframework.messaging.MessageHandler;
30+
import org.springframework.messaging.support.ChannelInterceptor;
31+
import org.springframework.messaging.support.MessageBuilder;
32+
33+
import static org.assertj.core.api.Assertions.assertThat;
34+
35+
/**
36+
* Verifies that {@link DefaultPollableMessageSource} invokes the full
37+
* {@link ChannelInterceptor} contract around poll lifecycles (GH-3131),
38+
* while keeping the legacy {@code preSend} behavior intact.
39+
*/
40+
class DefaultPollableMessageSourceTests {
41+
42+
private final List<String> events = new ArrayList<>();
43+
44+
private final AtomicInteger receiveCount = new AtomicInteger();
45+
46+
@Test
47+
void successfulPollInvokesFullInterceptorLifecycleInOrder() {
48+
DefaultPollableMessageSource source = newSource();
49+
source.setSource(this::message);
50+
source.addInterceptor(recorder());
51+
52+
assertThat(source.poll(noopHandler())).isTrue();
53+
54+
assertThat(this.events).containsExactly(
55+
"preReceive",
56+
"postReceive",
57+
"preSend",
58+
"afterReceiveCompletion",
59+
"postSend",
60+
"afterSendCompletion");
61+
}
62+
63+
@Test
64+
void falsePreReceiveShortCircuitsReceive() {
65+
DefaultPollableMessageSource source = newSource();
66+
source.setSource(this::countingMessage);
67+
source.addInterceptor(new Recorder() {
68+
@Override
69+
public boolean preReceive(MessageChannel channel) {
70+
DefaultPollableMessageSourceTests.this.events.add("preReceive");
71+
return false;
72+
}
73+
});
74+
75+
assertThat(source.poll(noopHandler())).isFalse();
76+
assertThat(this.events).containsExactly("preReceive");
77+
assertThat(this.receiveCount.get()).isZero();
78+
}
79+
80+
@Test
81+
void nullPostReceiveAbortsFurtherProcessing() {
82+
DefaultPollableMessageSource source = newSource();
83+
source.setSource(this::countingMessage);
84+
source.addInterceptor(new Recorder() {
85+
@Override
86+
public Message<?> postReceive(Message<?> message, MessageChannel channel) {
87+
DefaultPollableMessageSourceTests.this.events.add("postReceive");
88+
return null;
89+
}
90+
});
91+
92+
assertThat(source.poll(noopHandler())).isFalse();
93+
assertThat(this.receiveCount.get()).isEqualTo(1);
94+
assertThat(this.events).containsExactly(
95+
"preReceive",
96+
"postReceive",
97+
"afterReceiveCompletion");
98+
}
99+
100+
@Test
101+
void handlerFailureSkipsPostSendAndReportsExceptionOnCompletion() {
102+
DefaultPollableMessageSource source = newSource();
103+
DirectChannel errorChannel = new DirectChannel();
104+
errorChannel.subscribe(message -> {
105+
});
106+
source.setErrorChannel(errorChannel);
107+
source.setSource(this::message);
108+
source.addInterceptor(recorder());
109+
MessageHandler failingHandler = message -> {
110+
throw new IllegalStateException("boom");
111+
};
112+
113+
assertThat(source.poll(failingHandler)).isTrue();
114+
115+
assertThat(this.events).containsExactly(
116+
"preReceive",
117+
"postReceive",
118+
"preSend",
119+
"afterReceiveCompletion",
120+
"afterSendCompletion:exception");
121+
}
122+
123+
@Test
124+
void interceptorsAddedAfterSetSourceAreHonored() {
125+
DefaultPollableMessageSource source = newSource();
126+
source.setSource(this::message);
127+
source.addInterceptor(recorder());
128+
129+
assertThat(source.poll(noopHandler())).isTrue();
130+
131+
assertThat(this.events).contains("postSend", "afterSendCompletion");
132+
}
133+
134+
@Test
135+
void nullPreSendStillAbortsLegacyPath() {
136+
DefaultPollableMessageSource source = newSource();
137+
source.setSource(this::countingMessage);
138+
source.addInterceptor(new Recorder() {
139+
@Override
140+
public Message<?> preSend(Message<?> message, MessageChannel channel) {
141+
DefaultPollableMessageSourceTests.this.events.add("preSend");
142+
return null;
143+
}
144+
});
145+
146+
assertThat(source.poll(noopHandler())).isFalse();
147+
assertThat(this.receiveCount.get()).isEqualTo(1);
148+
assertThat(this.events).containsExactly(
149+
"preReceive",
150+
"postReceive",
151+
"preSend",
152+
"afterReceiveCompletion");
153+
}
154+
155+
private DefaultPollableMessageSource newSource() {
156+
return new DefaultPollableMessageSource(null);
157+
}
158+
159+
private Message<Object> message() {
160+
return MessageBuilder.withPayload((Object) "hello").build();
161+
}
162+
163+
private Message<Object> countingMessage() {
164+
this.receiveCount.incrementAndGet();
165+
return message();
166+
}
167+
168+
private ChannelInterceptor recorder() {
169+
return new Recorder();
170+
}
171+
172+
private MessageHandler noopHandler() {
173+
return message -> {
174+
};
175+
}
176+
177+
private class Recorder implements ChannelInterceptor {
178+
179+
@Override
180+
public boolean preReceive(MessageChannel channel) {
181+
DefaultPollableMessageSourceTests.this.events.add("preReceive");
182+
return true;
183+
}
184+
185+
@Override
186+
public Message<?> postReceive(Message<?> message, MessageChannel channel) {
187+
DefaultPollableMessageSourceTests.this.events.add("postReceive");
188+
return message;
189+
}
190+
191+
@Override
192+
public void afterReceiveCompletion(Message<?> message, MessageChannel channel,
193+
Exception ex) {
194+
record("afterReceiveCompletion", ex);
195+
}
196+
197+
@Override
198+
public Message<?> preSend(Message<?> message, MessageChannel channel) {
199+
DefaultPollableMessageSourceTests.this.events.add("preSend");
200+
return message;
201+
}
202+
203+
@Override
204+
public void postSend(Message<?> message, MessageChannel channel, boolean sent) {
205+
DefaultPollableMessageSourceTests.this.events.add("postSend");
206+
}
207+
208+
@Override
209+
public void afterSendCompletion(Message<?> message, MessageChannel channel,
210+
boolean sent, Exception ex) {
211+
record("afterSendCompletion", ex);
212+
}
213+
214+
private void record(String name, Exception ex) {
215+
DefaultPollableMessageSourceTests.this.events
216+
.add(ex != null ? name + ":exception" : name);
217+
}
218+
219+
}
220+
221+
}

0 commit comments

Comments
 (0)