diff --git a/spring-integration-mqtt/src/main/java/org/springframework/integration/mqtt/inbound/MqttPahoMessageDrivenChannelAdapter.java b/spring-integration-mqtt/src/main/java/org/springframework/integration/mqtt/inbound/MqttPahoMessageDrivenChannelAdapter.java index 8eb9af3116..3adcd19961 100644 --- a/spring-integration-mqtt/src/main/java/org/springframework/integration/mqtt/inbound/MqttPahoMessageDrivenChannelAdapter.java +++ b/spring-integration-mqtt/src/main/java/org/springframework/integration/mqtt/inbound/MqttPahoMessageDrivenChannelAdapter.java @@ -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); } diff --git a/spring-integration-mqtt/src/test/java/org/springframework/integration/mqtt/MqttAdapterTests.java b/spring-integration-mqtt/src/test/java/org/springframework/integration/mqtt/MqttAdapterTests.java index cbf46c13ae..e054d1cfdd 100644 --- a/spring-integration-mqtt/src/test/java/org/springframework/integration/mqtt/MqttAdapterTests.java +++ b/spring-integration-mqtt/src/test/java/org/springframework/integration/mqtt/MqttAdapterTests.java @@ -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() { - - @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() { - - @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() { - - @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("Hello, world!")); @@ -221,60 +205,43 @@ public class MqttAdapterTests { factory.setWill(will); factory = spy(factory); - final MqttAsyncClient client = mock(MqttAsyncClient.class); - doAnswer(new Answer() { + 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() { - - @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 callback = new AtomicReference(); - doAnswer(new Answer() { + 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 events = new LinkedBlockingQueue(); - doAnswer(new Answer() { - - @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()); } }