diff --git a/src/main/java/com/pusher/client/connection/websocket/WebSocketConnection.java b/src/main/java/com/pusher/client/connection/websocket/WebSocketConnection.java index 4a397d79..81f650ce 100644 --- a/src/main/java/com/pusher/client/connection/websocket/WebSocketConnection.java +++ b/src/main/java/com/pusher/client/connection/websocket/WebSocketConnection.java @@ -45,6 +45,7 @@ public class WebSocketConnection implements InternalConnection, WebSocketListene private final Consumer eventHandler; private String socketId; private int reconnectAttempts = 0; + private volatile Future reconnectTimer; public WebSocketConnection( final String url, @@ -93,7 +94,13 @@ private void tryConnecting() { @Override public void disconnect() { factory.queueOnEventThread(() -> { - if (canDisconnect()) { + if (state == ConnectionState.RECONNECTING) { + // The previous socket has already closed and won't call onClose, + // so finish disconnecting here instead of waiting for it. + underlyingConnection.removeWebSocketListener(); + updateState(ConnectionState.DISCONNECTING); + cancelTimeoutsAndTransitionToDisconnected(); + } else if (canDisconnect()) { updateState(ConnectionState.DISCONNECTING); underlyingConnection.close(); } @@ -253,15 +260,16 @@ private void tryReconnecting() { updateState(ConnectionState.RECONNECTING); long reconnectInterval = Math.min(maxReconnectionGap, reconnectAttempts * reconnectAttempts); - factory + reconnectTimer = factory .getTimers() .schedule( - () -> { + // Check the state on the event thread so it can't race with disconnect() + () -> factory.queueOnEventThread(() -> { if (state == ConnectionState.RECONNECTING) { underlyingConnection.removeWebSocketListener(); tryConnecting(); } - }, + }), reconnectInterval, TimeUnit.SECONDS ); @@ -275,6 +283,10 @@ private boolean shouldReconnect(int code) { private void cancelTimeoutsAndTransitionToDisconnected() { activityTimer.cancelTimeouts(); + if (reconnectTimer != null) { + reconnectTimer.cancel(false); + reconnectTimer = null; + } factory.queueOnEventThread(() -> { if (state == ConnectionState.DISCONNECTING) { diff --git a/src/test/java/com/pusher/client/connection/websocket/WebSocketConnectionTest.java b/src/test/java/com/pusher/client/connection/websocket/WebSocketConnectionTest.java index afee2cd6..b8106f9e 100644 --- a/src/test/java/com/pusher/client/connection/websocket/WebSocketConnectionTest.java +++ b/src/test/java/com/pusher/client/connection/websocket/WebSocketConnectionTest.java @@ -6,6 +6,7 @@ import static org.mockito.Mockito.any; import static org.mockito.Mockito.anyString; import static org.mockito.Mockito.doAnswer; +import static org.mockito.Mockito.doReturn; import static org.mockito.Mockito.doThrow; import static org.mockito.Mockito.mock; import static org.mockito.Mockito.never; @@ -24,6 +25,7 @@ import org.junit.Before; import org.junit.Test; import org.junit.runner.RunWith; +import org.mockito.ArgumentCaptor; import org.mockito.Mock; import org.mockito.invocation.InvocationOnMock; import org.mockito.runners.MockitoJUnitRunner; @@ -34,6 +36,7 @@ import java.net.URI; import java.net.URISyntaxException; import java.util.concurrent.ScheduledExecutorService; +import java.util.concurrent.ScheduledFuture; import java.util.concurrent.ScheduledThreadPoolExecutor; import java.util.concurrent.TimeUnit; import java.util.function.Consumer; @@ -466,6 +469,57 @@ public Object answer(InvocationOnMock invocation) { assertEquals(ConnectionState.DISCONNECTED, connection.getState()); } + @Test + public void testDisconnectInReconnectingStateTransitionsToDisconnected() { + connect(); + connection.onClose(3999, "reason", true); + assertEquals(ConnectionState.RECONNECTING, connection.getState()); + + connection.disconnect(); + + assertEquals(ConnectionState.DISCONNECTED, connection.getState()); + verify(mockEventListener) + .onConnectionStateChange(new ConnectionStateChange(ConnectionState.RECONNECTING, ConnectionState.DISCONNECTING)); + verify(mockEventListener) + .onConnectionStateChange(new ConnectionStateChange(ConnectionState.DISCONNECTING, ConnectionState.DISCONNECTED)); + verify(mockUnderlyingConnection, never()).close(); + verify(factory).shutdownThreads(); + } + + @Test + public void testDisconnectInReconnectingStateCancelsPendingReconnect() { + final ScheduledFuture reconnectFuture = mock(ScheduledFuture.class); + when(factory.getTimers()).thenReturn(scheduledExecutorService); + doReturn(reconnectFuture) + .when(scheduledExecutorService) + .schedule(any(Runnable.class), any(Long.class), any(TimeUnit.class)); + + connection.connect(); + connection.onClose(500, "reason", true); + assertEquals(ConnectionState.RECONNECTING, connection.getState()); + + connection.disconnect(); + + verify(reconnectFuture).cancel(false); + } + + @Test + public void testReconnectThatRunsAfterDisconnectDoesNotReconnect() throws SSLException { + final ArgumentCaptor reconnect = ArgumentCaptor.forClass(Runnable.class); + when(factory.getTimers()).thenReturn(scheduledExecutorService); + + connection.connect(); + connection.onClose(500, "reason", true); + verify(scheduledExecutorService).schedule(reconnect.capture(), any(Long.class), any(TimeUnit.class)); + + connection.disconnect(); + // Simulate the reconnect task already running when disconnect() cancelled it + reconnect.getValue().run(); + + verify(factory, times(1)).newWebSocketClientWrapper(any(URI.class), any(Proxy.class), any(WebSocketConnection.class)); + assertEquals(ConnectionState.DISCONNECTED, connection.getState()); + } + /* end of tests */ private void connect() {