diff --git a/src/main/java/net/spy/memcached/CacheManager.java b/src/main/java/net/spy/memcached/CacheManager.java index 7b327cddf..7703facc8 100644 --- a/src/main/java/net/spy/memcached/CacheManager.java +++ b/src/main/java/net/spy/memcached/CacheManager.java @@ -23,7 +23,6 @@ import java.text.SimpleDateFormat; import java.util.AbstractMap; import java.util.ArrayList; -import java.util.Collections; import java.util.Date; import java.util.List; import java.util.Map; @@ -718,7 +717,7 @@ public void connectionEstablished(MemcachedNode node, int reconnectCount) { } } }; - cfb.setInitialObservers(Collections.singleton(observer)); + cfb.addInitialObserver(observer); int poolId = CacheManager.POOL_ID.getAndIncrement(); client = new ArcusClient[poolSize]; @@ -740,6 +739,8 @@ public void connectionEstablished(MemcachedNode node, int reconnectCount) { } client = null; return; + } finally { + cfb.removeInitialObserver(observer); } try { diff --git a/src/main/java/net/spy/memcached/ConnectionFactoryBuilder.java b/src/main/java/net/spy/memcached/ConnectionFactoryBuilder.java index 2cb5b13b7..66ebb528d 100644 --- a/src/main/java/net/spy/memcached/ConnectionFactoryBuilder.java +++ b/src/main/java/net/spy/memcached/ConnectionFactoryBuilder.java @@ -19,8 +19,8 @@ import java.io.IOException; import java.net.InetSocketAddress; +import java.util.ArrayList; import java.util.Collection; -import java.util.Collections; import java.util.HashMap; import java.util.List; import java.util.Map; @@ -51,8 +51,7 @@ public class ConnectionFactoryBuilder { private FailureMode failureMode = FailureMode.Cancel; - private Collection initialObservers - = Collections.emptyList(); + private List initialObservers = new ArrayList<>(); private OperationFactory opFact; @@ -202,16 +201,32 @@ public ConnectionFactoryBuilder setFailureMode(FailureMode fm) { /** * Set the initial connection observers (will observe initial connection). */ - public ConnectionFactoryBuilder setInitialObservers( - Collection obs) { + public ConnectionFactoryBuilder setInitialObservers(Collection obs) { if (obs == null || obs.isEmpty()) { throw new IllegalArgumentException("Initial observers must not be null or empty."); } - initialObservers = obs; + initialObservers.clear(); + initialObservers.addAll(obs); return this; } + void addInitialObserver(ConnectionObserver observer) { + if (observer == null) { + throw new IllegalArgumentException("Initial observer must not be null."); + } + + initialObservers.add(observer); + } + + void removeInitialObserver(ConnectionObserver observer) { + if (observer == null) { + throw new IllegalArgumentException("Initial observer must not be null."); + } + + initialObservers.remove(observer); + } + /** * Set the operation factory. * diff --git a/src/test/java/net/spy/memcached/ArcusClientInitialObserverTest.java b/src/test/java/net/spy/memcached/ArcusClientInitialObserverTest.java new file mode 100644 index 000000000..9f183234d --- /dev/null +++ b/src/test/java/net/spy/memcached/ArcusClientInitialObserverTest.java @@ -0,0 +1,63 @@ +package net.spy.memcached; + +import java.util.Collection; +import java.util.Collections; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.TimeUnit; + +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.Test; + +import static org.junit.jupiter.api.Assertions.assertAll; +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertTrue; + +class ArcusClientInitialObserverTest { + + private ArcusClient client; + + @AfterEach + void tearDown() { + if (client != null) { + client.shutdown(); + } + } + + @Test + void test() { + // given + CountDownLatch latch = new CountDownLatch(1); + + ConnectionObserver observer = new ConnectionObserver() { + @Override + public void connectionEstablished(MemcachedNode node, int reconnectCount) { + latch.countDown(); + } + + @Override + public void connectionLost(MemcachedNode node) { + // do-nothing. + } + }; + + // when + ConnectionFactoryBuilder cfb = new ConnectionFactoryBuilder() + .setInitialObservers(Collections.singletonList(observer)); + + client = ArcusClient.createArcusClient( + "127.0.0.1:2181", + "test", + cfb + ); + + Collection configuredObservers = cfb.build().getInitialObservers(); + + // then + assertAll( + () -> assertTrue(latch.await(700, TimeUnit.MILLISECONDS)), + () -> assertTrue(client.removeObserver(observer)), + () -> assertEquals(1, configuredObservers.size()), + () -> assertTrue(configuredObservers.contains(observer)) + ); + } +}