AMQP-4238: Detect Subscription Failures and QOS

JIRA: https://jira.spring.io/browse/INT-4238

Revert to using the sync client in the message-driven adapter so we can detect
subscription failures (the sync client throws an exception).

The only reason to use the async client was to timeout disconnects; this can
be achieved with the sync client and `disconnectForcibly`.

Also, the subscribe method updates the qos argument with the granted QOS values.

Detect and log if any QOS does not match the request.

Polishing

Polishing - PR Comments

Conflicts:
	spring-integration-mqtt/src/test/java/org/springframework/integration/mqtt/MqttAdapterTests.java

* Remove all new tests since Paho lib has class signature check, so we can't mock its classes
This commit is contained in:
Gary Russell
2017-03-06 14:34:10 -05:00
committed by Artem Bilan
parent 0070d38d0c
commit 789ed9d1b8
2 changed files with 100 additions and 131 deletions

View File

@@ -20,9 +20,10 @@ import java.util.Arrays;
import java.util.Date;
import java.util.concurrent.ScheduledFuture;
import org.eclipse.paho.client.mqttv3.IMqttAsyncClient;
import org.eclipse.paho.client.mqttv3.IMqttClient;
import org.eclipse.paho.client.mqttv3.IMqttDeliveryToken;
import org.eclipse.paho.client.mqttv3.MqttCallback;
import org.eclipse.paho.client.mqttv3.MqttClient;
import org.eclipse.paho.client.mqttv3.MqttConnectOptions;
import org.eclipse.paho.client.mqttv3.MqttException;
import org.eclipse.paho.client.mqttv3.MqttMessage;
@@ -54,7 +55,7 @@ public class MqttPahoMessageDrivenChannelAdapter extends AbstractMqttMessageDriv
private final MqttPahoClientFactory clientFactory;
private volatile IMqttAsyncClient client;
private volatile IMqttClient client;
private volatile ScheduledFuture<?> reconnectFuture;
@@ -110,7 +111,7 @@ public class MqttPahoMessageDrivenChannelAdapter extends AbstractMqttMessageDriv
}
/**
* Set the completion timeout for async operations. Not settable using the namespace.
* Set the completion timeout for operations. Not settable using the namespace.
* Default 30000 milliseconds.
* @param completionTimeout The timeout.
* @since 4.1
@@ -159,16 +160,14 @@ public class MqttPahoMessageDrivenChannelAdapter extends AbstractMqttMessageDriv
if (this.consumerStopAction.equals(ConsumerStopAction.UNSUBSCRIBE_ALWAYS)
|| (this.consumerStopAction.equals(ConsumerStopAction.UNSUBSCRIBE_CLEAN)
&& this.cleanSession)) {
this.client.unsubscribe(getTopic())
.waitForCompletion(this.completionTimeout);
this.client.unsubscribe(getTopic());
}
}
catch (MqttException e) {
logger.error("Exception while unsubscribing", e);
}
try {
this.client.disconnect()
.waitForCompletion(this.completionTimeout);
this.client.disconnectForcibly(this.completionTimeout);
}
catch (MqttException e) {
logger.error("Exception while disconnecting", e);
@@ -190,8 +189,7 @@ public class MqttPahoMessageDrivenChannelAdapter extends AbstractMqttMessageDriv
try {
super.addTopic(topic, qos);
if (this.client != null && this.client.isConnected()) {
this.client.subscribe(topic, qos)
.waitForCompletion(this.completionTimeout);
this.client.subscribe(topic, qos);
}
}
catch (MqttException e) {
@@ -208,8 +206,7 @@ public class MqttPahoMessageDrivenChannelAdapter extends AbstractMqttMessageDriv
this.topicLock.lock();
try {
if (this.client != null && this.client.isConnected()) {
this.client.unsubscribe(topic)
.waitForCompletion(this.completionTimeout);
this.client.unsubscribe(topic);
}
super.removeTopic(topic);
}
@@ -230,23 +227,36 @@ public class MqttPahoMessageDrivenChannelAdapter extends AbstractMqttMessageDriv
}
Assert.state(getUrl() != null || connectionOptions.getServerURIs() != null,
"If no 'url' provided, connectionOptions.getServerURIs() must not be null");
this.client = this.clientFactory.getAsyncClientInstance(getUrl(), getClientId());
this.client = this.clientFactory.getClientInstance(getUrl(), getClientId());
this.client.setCallback(this);
if (this.client instanceof MqttClient) {
((MqttClient) this.client).setTimeToWait(this.completionTimeout);
}
this.topicLock.lock();
String[] topics = getTopic();
try {
this.client.connect(connectionOptions)
.waitForCompletion(this.completionTimeout);
this.client.subscribe(getTopic(), getQos())
.waitForCompletion(this.completionTimeout);
this.client.connect(connectionOptions);
int[] requestedQos = getQos();
int[] grantedQos = Arrays.copyOf(requestedQos, requestedQos.length);
this.client.subscribe(topics, grantedQos);
for (int i = 0; i < requestedQos.length; i++) {
if (grantedQos[i] != requestedQos[i]) {
if (logger.isWarnEnabled()) {
logger.warn("Granted QOS different to Requested QOS; topics: " + Arrays.toString(topics)
+ " requested: " + Arrays.toString(requestedQos)
+ " granted: " + Arrays.toString(grantedQos));
}
break;
}
}
}
catch (MqttException e) {
if (this.applicationEventPublisher != null) {
this.applicationEventPublisher.publishEvent(new MqttConnectionFailedEvent(this, e));
}
logger.error("Error connecting or subscribing to " + Arrays.asList(getTopic()), e);
this.client.disconnect()
.waitForCompletion(this.completionTimeout);
logger.error("Error connecting or subscribing to " + Arrays.toString(topics), e);
this.client.disconnectForcibly(this.completionTimeout);
throw e;
}
finally {
@@ -254,7 +264,7 @@ public class MqttPahoMessageDrivenChannelAdapter extends AbstractMqttMessageDriv
}
if (this.client.isConnected()) {
this.connected = true;
String message = "Connected and subscribed to " + Arrays.asList(getTopic());
String message = "Connected and subscribed to " + Arrays.toString(topics);
if (logger.isDebugEnabled()) {
logger.debug(message);
}

View File

@@ -25,16 +25,15 @@ import static org.junit.Assert.assertThat;
import static org.junit.Assert.assertTrue;
import static org.mockito.BDDMockito.given;
import static org.mockito.BDDMockito.willAnswer;
import static org.mockito.Matchers.any;
import static org.mockito.Matchers.anyString;
import static org.mockito.Mockito.doAnswer;
import static org.mockito.Mockito.doReturn;
import static org.mockito.BDDMockito.willReturn;
import static org.mockito.Mockito.any;
import static org.mockito.Mockito.anyLong;
import static org.mockito.Mockito.anyString;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.never;
import static org.mockito.Mockito.spy;
import static org.mockito.Mockito.times;
import static org.mockito.Mockito.verify;
import static org.mockito.Mockito.when;
import java.util.Properties;
import java.util.concurrent.BlockingQueue;
@@ -49,6 +48,7 @@ import javax.net.SocketFactory;
import org.aopalliance.intercept.MethodInterceptor;
import org.apache.commons.logging.Log;
import org.eclipse.paho.client.mqttv3.IMqttClient;
import org.eclipse.paho.client.mqttv3.IMqttToken;
import org.eclipse.paho.client.mqttv3.MqttAsyncClient;
import org.eclipse.paho.client.mqttv3.MqttCallback;
@@ -60,8 +60,6 @@ import org.eclipse.paho.client.mqttv3.MqttSecurityException;
import org.eclipse.paho.client.mqttv3.MqttToken;
import org.eclipse.paho.client.mqttv3.persist.MemoryPersistence;
import org.junit.Test;
import org.mockito.invocation.InvocationOnMock;
import org.mockito.stubbing.Answer;
import org.springframework.aop.framework.ProxyFactoryBean;
import org.springframework.beans.DirectFieldAccessor;
@@ -147,13 +145,7 @@ public class MqttAdapterTests {
factory = spy(factory);
final MqttAsyncClient client = mock(MqttAsyncClient.class);
doAnswer(new Answer<MqttAsyncClient>() {
@Override
public MqttAsyncClient answer(InvocationOnMock invocation) throws Throwable {
return client;
}
}).when(factory).getAsyncClientInstance(anyString(), anyString());
willAnswer(invocation -> client).given(factory).getAsyncClientInstance(anyString(), anyString());
MqttPahoMessageHandler handler = new MqttPahoMessageHandler("foo", "bar", factory);
handler.setDefaultTopic("mqtt-foo");
@@ -163,39 +155,31 @@ public class MqttAdapterTests {
final MqttToken token = mock(MqttToken.class);
final AtomicBoolean connectCalled = new AtomicBoolean();
doAnswer(new Answer<MqttToken>() {
@Override
public MqttToken answer(InvocationOnMock invocation) throws Throwable {
MqttConnectOptions options = (MqttConnectOptions) invocation.getArguments()[0];
assertEquals(23, options.getConnectionTimeout());
assertEquals(45, options.getKeepAliveInterval());
assertEquals("pass", new String(options.getPassword()));
assertSame(socketFactory, options.getSocketFactory());
assertSame(props, options.getSSLProperties());
assertEquals("user", options.getUserName());
assertEquals("foo", options.getWillDestination());
assertEquals("bar", new String(options.getWillMessage().getPayload()));
assertEquals(2, options.getWillMessage().getQos());
connectCalled.set(true);
return token;
}
}).when(client).connect(any(MqttConnectOptions.class));
doReturn(token).when(client).subscribe(any(String[].class), any(int[].class));
willAnswer(invocation -> {
MqttConnectOptions options = (MqttConnectOptions) invocation.getArguments()[0];
assertEquals(23, options.getConnectionTimeout());
assertEquals(45, options.getKeepAliveInterval());
assertEquals("pass", new String(options.getPassword()));
assertSame(socketFactory, options.getSocketFactory());
assertSame(props, options.getSSLProperties());
assertEquals("user", options.getUserName());
assertEquals("foo", options.getWillDestination());
assertEquals("bar", new String(options.getWillMessage().getPayload()));
assertEquals(2, options.getWillMessage().getQos());
connectCalled.set(true);
return token;
}).given(client).connect(any(MqttConnectOptions.class));
willReturn(token).given(client).subscribe(any(String[].class), any(int[].class));
final MqttDeliveryToken deliveryToken = mock(MqttDeliveryToken.class);
final AtomicBoolean publishCalled = new AtomicBoolean();
doAnswer(new Answer<MqttDeliveryToken>() {
@Override
public MqttDeliveryToken answer(InvocationOnMock invocation) throws Throwable {
assertEquals("mqtt-foo", invocation.getArguments()[0]);
MqttMessage message = (MqttMessage) invocation.getArguments()[1];
assertEquals("Hello, world!", new String(message.getPayload()));
publishCalled.set(true);
return deliveryToken;
}
}).when(client).publish(anyString(), any(MqttMessage.class));
willAnswer(invocation -> {
assertEquals("mqtt-foo", invocation.getArguments()[0]);
MqttMessage message = (MqttMessage) invocation.getArguments()[1];
assertEquals("Hello, world!", new String(message.getPayload()));
publishCalled.set(true);
return deliveryToken;
}).given(client).publish(anyString(), any(MqttMessage.class));
handler.handleMessage(new GenericMessage<String>("Hello, world!"));
@@ -221,60 +205,43 @@ public class MqttAdapterTests {
factory.setWill(will);
factory = spy(factory);
final MqttAsyncClient client = mock(MqttAsyncClient.class);
doAnswer(new Answer<MqttAsyncClient>() {
final IMqttClient client = mock(IMqttClient.class);
willAnswer(invocation -> client).given(factory).getClientInstance(anyString(), anyString());
@Override
public MqttAsyncClient answer(InvocationOnMock invocation) throws Throwable {
return client;
}
}).when(factory).getAsyncClientInstance(anyString(), anyString());
final MqttToken token = mock(MqttToken.class);
final AtomicBoolean connectCalled = new AtomicBoolean();
final AtomicBoolean failConnection = new AtomicBoolean();
final CountDownLatch waitToFail = new CountDownLatch(1);
final CountDownLatch failInProcess = new CountDownLatch(1);
final CountDownLatch goodConnection = new CountDownLatch(2);
final MqttException reconnectException = new MqttException(MqttException.REASON_CODE_SERVER_CONNECT_ERROR);
doAnswer(new Answer<Object>() {
@Override
public Object answer(InvocationOnMock invocation) throws Throwable {
if (failConnection.get()) {
failInProcess.countDown();
waitToFail.await(10, TimeUnit.SECONDS);
throw reconnectException;
}
MqttConnectOptions options = (MqttConnectOptions) invocation.getArguments()[0];
assertEquals(23, options.getConnectionTimeout());
assertEquals(45, options.getKeepAliveInterval());
assertEquals("pass", new String(options.getPassword()));
assertSame(socketFactory, options.getSocketFactory());
assertSame(props, options.getSSLProperties());
assertEquals("user", options.getUserName());
assertEquals("foo", options.getWillDestination());
assertEquals("bar", new String(options.getWillMessage().getPayload()));
assertEquals(2, options.getWillMessage().getQos());
connectCalled.set(true);
goodConnection.countDown();
return token;
willAnswer(invocation -> {
if (failConnection.get()) {
failInProcess.countDown();
waitToFail.await(10, TimeUnit.SECONDS);
throw reconnectException;
}
}).when(client).connect(any(MqttConnectOptions.class));
doReturn(token).when(client).subscribe(any(String[].class), any(int[].class));
doReturn(token).when(client).disconnect();
MqttConnectOptions options = (MqttConnectOptions) invocation.getArguments()[0];
assertEquals(23, options.getConnectionTimeout());
assertEquals(45, options.getKeepAliveInterval());
assertEquals("pass", new String(options.getPassword()));
assertSame(socketFactory, options.getSocketFactory());
assertSame(props, options.getSSLProperties());
assertEquals("user", options.getUserName());
assertEquals("foo", options.getWillDestination());
assertEquals("bar", new String(options.getWillMessage().getPayload()));
assertEquals(2, options.getWillMessage().getQos());
connectCalled.set(true);
goodConnection.countDown();
return null;
}).given(client).connect(any(MqttConnectOptions.class));
final AtomicReference<MqttCallback> callback = new AtomicReference<MqttCallback>();
doAnswer(new Answer<Object>() {
willAnswer(invocation -> {
callback.set((MqttCallback) invocation.getArguments()[0]);
return null;
}).given(client).setCallback(any(MqttCallback.class));
@Override
public Object answer(InvocationOnMock invocation) throws Throwable {
callback.set((MqttCallback) invocation.getArguments()[0]);
return null;
}
}).when(client).setCallback(any(MqttCallback.class));
when(client.isConnected()).thenReturn(true);
given(client.isConnected()).willReturn(true);
MqttPahoMessageDrivenChannelAdapter adapter = new MqttPahoMessageDrivenChannelAdapter("foo", "bar", factory,
"baz", "fix");
@@ -286,14 +253,10 @@ public class MqttAdapterTests {
adapter.setBeanFactory(mock(BeanFactory.class));
ApplicationEventPublisher applicationEventPublisher = mock(ApplicationEventPublisher.class);
final BlockingQueue<MqttIntegrationEvent> events = new LinkedBlockingQueue<MqttIntegrationEvent>();
doAnswer(new Answer<Void>() {
@Override
public Void answer(InvocationOnMock invocation) throws Throwable {
events.add((MqttIntegrationEvent) invocation.getArguments()[0]);
return null;
}
}).when(applicationEventPublisher).publishEvent(any(MqttIntegrationEvent.class));
willAnswer(invocation -> {
events.add((MqttIntegrationEvent) invocation.getArguments()[0]);
return null;
}).given(applicationEventPublisher).publishEvent(any(MqttIntegrationEvent.class));
adapter.setApplicationEventPublisher(applicationEventPublisher);
adapter.setRecoveryInterval(500);
adapter.afterPropertiesSet();
@@ -341,7 +304,7 @@ public class MqttAdapterTests {
@Test
public void testStopActionDefault() throws Exception {
final MqttAsyncClient client = mock(MqttAsyncClient.class);
final IMqttClient client = mock(IMqttClient.class);
MqttPahoMessageDrivenChannelAdapter adapter = buildAdapter(client, null, null);
adapter.start();
@@ -351,7 +314,7 @@ public class MqttAdapterTests {
@Test
public void testStopActionDefaultNotClean() throws Exception {
final MqttAsyncClient client = mock(MqttAsyncClient.class);
final IMqttClient client = mock(IMqttClient.class);
MqttPahoMessageDrivenChannelAdapter adapter = buildAdapter(client, false, null);
adapter.start();
@@ -361,7 +324,7 @@ public class MqttAdapterTests {
@Test
public void testStopActionAlways() throws Exception {
final MqttAsyncClient client = mock(MqttAsyncClient.class);
final IMqttClient client = mock(IMqttClient.class);
MqttPahoMessageDrivenChannelAdapter adapter = buildAdapter(client, false,
ConsumerStopAction.UNSUBSCRIBE_ALWAYS);
@@ -372,7 +335,7 @@ public class MqttAdapterTests {
@Test
public void testStopActionNever() throws Exception {
final MqttAsyncClient client = mock(MqttAsyncClient.class);
final IMqttClient client = mock(IMqttClient.class);
MqttPahoMessageDrivenChannelAdapter adapter = buildAdapter(client, null, ConsumerStopAction.UNSUBSCRIBE_NEVER);
adapter.start();
@@ -382,7 +345,7 @@ public class MqttAdapterTests {
@Test
public void testReconnect() throws Exception {
final MqttAsyncClient client = mock(MqttAsyncClient.class);
final IMqttClient client = mock(IMqttClient.class);
MqttPahoMessageDrivenChannelAdapter adapter = buildAdapter(client, null, ConsumerStopAction.UNSUBSCRIBE_NEVER);
adapter.setRecoveryInterval(10);
Log logger = spy(TestUtils.getPropertyValue(adapter, "logger", Log.class));
@@ -408,12 +371,12 @@ public class MqttAdapterTests {
taskScheduler.destroy();
}
private MqttPahoMessageDrivenChannelAdapter buildAdapter(final MqttAsyncClient client, Boolean cleanSession,
private MqttPahoMessageDrivenChannelAdapter buildAdapter(final IMqttClient client, Boolean cleanSession,
ConsumerStopAction action) throws MqttException, MqttSecurityException {
DefaultMqttPahoClientFactory factory = new DefaultMqttPahoClientFactory() {
@Override
public MqttAsyncClient getAsyncClientInstance(String uri, String clientId) throws MqttException {
public IMqttClient getClientInstance(String uri, String clientId) throws MqttException {
return client;
}
@@ -425,11 +388,7 @@ public class MqttAdapterTests {
if (action != null) {
factory.setConsumerStopAction(action);
}
when(client.connect(any(MqttConnectOptions.class))).thenReturn(this.alwaysComplete);
when(client.subscribe(any(String[].class), any(int[].class))).thenReturn(this.alwaysComplete);
when(client.disconnect()).thenReturn(this.alwaysComplete);
when(client.unsubscribe(any(String[].class))).thenReturn(this.alwaysComplete);
when(client.isConnected()).thenReturn(true);
given(client.isConnected()).willReturn(true);
MqttPahoMessageDrivenChannelAdapter adapter = new MqttPahoMessageDrivenChannelAdapter("client", factory, "foo");
adapter.setApplicationEventPublisher(mock(ApplicationEventPublisher.class));
adapter.setOutputChannel(new NullChannel());
@@ -438,18 +397,18 @@ public class MqttAdapterTests {
return adapter;
}
private void verifyUnsubscribe(MqttAsyncClient client) throws Exception {
private void verifyUnsubscribe(IMqttClient client) throws Exception {
verify(client).connect(any(MqttConnectOptions.class));
verify(client).subscribe(any(String[].class), any(int[].class));
verify(client).unsubscribe(any(String[].class));
verify(client).disconnect();
verify(client).disconnectForcibly(anyLong());
}
private void verifyNotUnsubscribe(MqttAsyncClient client) throws Exception {
private void verifyNotUnsubscribe(IMqttClient client) throws Exception {
verify(client).connect(any(MqttConnectOptions.class));
verify(client).subscribe(any(String[].class), any(int[].class));
verify(client, never()).unsubscribe(any(String[].class));
verify(client).disconnect();
verify(client).disconnectForcibly(anyLong());
}
}