From 411aa296a446e5324ab4b4f99bac504a26cb1de1 Mon Sep 17 00:00:00 2001 From: Gary Russell Date: Wed, 29 Feb 2012 17:38:59 -0500 Subject: [PATCH] INT-1871 TCP CachingClientConnectionFactory Pool based on algorithm used for spring-integration-file CachingSessionFactory introduced by INT-2146. Refactored that code to use the common SimplePool. One difference to the previous implementation is the ability to change the pool size dynamically. If the size is reduced and more than the new size are in use, items are closed as they are returned until the pool size is as requested. Initial commit. Allow Pool Size Changes Factor out Pool Polishing Pool Tests Default forever Javadocs, File Polishing INT-1871 PR Polishing * Consistent/cleaner method names * Track checkouts; reject release of 'foreign' objects. * Add 'getAllocatedCount()' --- .../integration/util/Pool.java | 74 ++++ .../integration/util/SimplePool.java | 263 ++++++++++++++ .../integration/util/SimplePoolTests.java | 152 +++++++++ .../remote/session/CachingSessionFactory.java | 93 ++--- .../ftp/session/SessionFactoryTests.java | 3 +- .../CachingClientConnectionFactory.java | 323 ++++++++++++++++++ ...ngClientConnectionFactoryTests-context.xml | 46 +++ .../CachingClientConnectionFactoryTests.java | 293 ++++++++++++++++ 8 files changed, 1190 insertions(+), 57 deletions(-) create mode 100644 spring-integration-core/src/main/java/org/springframework/integration/util/Pool.java create mode 100644 spring-integration-core/src/main/java/org/springframework/integration/util/SimplePool.java create mode 100644 spring-integration-core/src/test/java/org/springframework/integration/util/SimplePoolTests.java create mode 100644 spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/connection/CachingClientConnectionFactory.java create mode 100644 spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/connection/CachingClientConnectionFactoryTests-context.xml create mode 100644 spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/connection/CachingClientConnectionFactoryTests.java diff --git a/spring-integration-core/src/main/java/org/springframework/integration/util/Pool.java b/spring-integration-core/src/main/java/org/springframework/integration/util/Pool.java new file mode 100644 index 0000000000..843475926b --- /dev/null +++ b/spring-integration-core/src/main/java/org/springframework/integration/util/Pool.java @@ -0,0 +1,74 @@ +/* + * Copyright 2002-2012 the original author or authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.springframework.integration.util; + +/** + * Represents a pool of items. + * @author Gary Russell + * @since 2.2 + * + */ +public interface Pool { + + /** + * Obtains an item from the pool. + * @return the item. + */ + T getItem(); + + /** + * Releases an item back into the pool. This must be an item that + * was previously retrieved using {@link #getItem()}. + * @param t the item. + * @throws IllegalArgumentException when a "foreign" object + * is released. + */ + void releaseItem(T t); + + /** + * Removes all idle items from the pool. + */ + void removeAllIdleItems(); + + /** + * Returns the current size (limit) of the pool. + * @return the size. + */ + int getPoolSize(); + + /** + * Returns the number of items that have been allocated + * but are not currently in use. + * @return The number of items. + */ + int getIdleCount(); + + /** + * Returns the number of allocated items that are currently + * checked out of the pool. + * @return The number of items. + */ + int getActiveCount(); + + /** + * Returns the current count of allocated items (in use and + * idle). May be less than the pool size, and reflects the + * high water mark of pool usage. + * @return the number of items. + */ + int getAllocatedCount(); + +} diff --git a/spring-integration-core/src/main/java/org/springframework/integration/util/SimplePool.java b/spring-integration-core/src/main/java/org/springframework/integration/util/SimplePool.java new file mode 100644 index 0000000000..3ce434d996 --- /dev/null +++ b/spring-integration-core/src/main/java/org/springframework/integration/util/SimplePool.java @@ -0,0 +1,263 @@ +/* + * Copyright 2002-2012 the original author or authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.springframework.integration.util; + +import java.util.Collections; +import java.util.HashSet; +import java.util.Set; +import java.util.concurrent.BlockingQueue; +import java.util.concurrent.LinkedBlockingQueue; +import java.util.concurrent.Semaphore; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicInteger; + +import org.apache.commons.logging.Log; +import org.apache.commons.logging.LogFactory; +import org.springframework.integration.MessagingException; +import org.springframework.util.Assert; + +/** + * Implementation of {@link Pool} supporting dynamic resizing and a variable + * timeout when attempting to obtain an item from the pool. Pool grows on + * demand up to the limit. + * @author Gary Russell + * @since 2.2 + * + */ +public class SimplePool implements Pool { + + protected final Log logger = LogFactory.getLog(this.getClass()); + + private final Semaphore permits = new Semaphore(0); + + private final AtomicInteger poolSize = new AtomicInteger(); + + private final AtomicInteger targetPoolSize = new AtomicInteger(); + + private long waitTimeout = Long.MAX_VALUE; + + private final BlockingQueue available = new LinkedBlockingQueue(); + + private final Set allocated = Collections.synchronizedSet(new HashSet()); + + private final PoolItemCallback callback; + + /** + * Creates a SimplePool with a specific limit. + * @param poolSize The maximum number of items the pool supports. + * @param callback A {@link PoolItemCallback} implementation called during various + * pool operations. + */ + public SimplePool(int poolSize, PoolItemCallback callback) { + if (poolSize <= 0) { + this.poolSize.set(Integer.MAX_VALUE); + this.targetPoolSize.set(Integer.MAX_VALUE); + this.permits.release(Integer.MAX_VALUE); + } + else { + this.poolSize.set(poolSize); + this.targetPoolSize.set(poolSize); + this.permits.release(poolSize); + } + this.callback = callback; + } + + /** + * Adjusts the current pool size. When reducing the pool size, attempts to + * remove the delta from the pool. If there are not enough unused items in + * the pool, the actual pool size will decrease to the specified size as in-use + * items are returned. + * @param poolSize The desired target pool size. + */ + public synchronized void setPoolSize(int poolSize) { + int delta = poolSize - this.poolSize.get(); + this.targetPoolSize.addAndGet(delta); + if (logger.isDebugEnabled()) { + logger.debug(String.format("Target pool size changed by %d, now %d", delta, this.targetPoolSize.get())); + } + if (delta > 0) { + this.poolSize.addAndGet(delta); + this.permits.release(delta); + } + else while (delta < 0) { + if (!this.permits.tryAcquire()) { + break; + } + T item = this.available.poll(); + if (item == null) { + this.permits.release(); + break; + } + doRemoveItem(item); + this.poolSize.decrementAndGet(); + delta++; + } + if (delta < 0 && logger.isDebugEnabled()) { + logger.debug(String.format("Pool is overcommitted by %d; items will be removed when returned", -delta)); + } + } + + /** + * Returns the current size of the pool; may be greater than the target pool size + * if it was recently reduced and too many items were in use to allow the new size + * to be set. + */ + public int getPoolSize() { + return this.poolSize.get(); + } + + public int getIdleCount() { + return this.available.size(); + } + + public int getActiveCount() { + return this.getAllocatedCount() - this.getIdleCount(); + } + + public int getAllocatedCount() { + return this.allocated.size(); + } + + /** + * Adjusts the wait timeout - the time for which getItem() will wait if no idle + * entries are available.
Default: infinity. + * @param waitTimeout The wait timeout in milliseconds. + */ + public void setWaitTimeout(long waitTimeout) { + this.waitTimeout = waitTimeout; + } + + /** + * Obtains an item from the pool; waits up to waitTime milliseconds (default infinity). + * @throws MessagingException if no items become available in time. + */ + public T getItem() { + boolean permitted = false; + try { + try { + permitted = this.permits.tryAcquire(this.waitTimeout, TimeUnit.MILLISECONDS); + } + catch (InterruptedException e) { + Thread.currentThread().interrupt(); + throw new MessagingException("Interrupted awaiting a pooled resource", e); + } + if (!permitted) { + throw new IllegalStateException("Timed out while waiting to aquire a pool entry."); + } + T item = doGetItem(); + return item; + } + catch (Exception e) { + if (permitted) { + this.permits.release(); + } + if (e instanceof MessagingException) { + throw (MessagingException) e; + } + throw new MessagingException("Failed to obtain pooled item", e); + } + } + + private T doGetItem() { + T item = this.available.poll(); + if (item != null && logger.isDebugEnabled()) { + logger.debug("Obtained " + item + " from pool."); + } + if (item == null) { + item = this.callback.createForPool(); + if (logger.isDebugEnabled()) { + logger.debug("Obtained new " + item + "."); + } + allocated.add(item); + } + else if (this.callback.isStale(item)) { + if (logger.isDebugEnabled()) { + logger.debug("Received a stale item, will attempt to get a new one."); + } + item = doGetItem(); + } + return item; + } + + /** + * Returns an item to the pool. Item may be null, in which case a subsequent getItem() + * will return a new instance. + */ + public synchronized void releaseItem(T item) { + Assert.isTrue(item == null || this.allocated.contains(item), + "You can only release items that were obtained from the pool"); + if (this.poolSize.get() > targetPoolSize.get()) { + poolSize.decrementAndGet(); + if (item != null) { + doRemoveItem(item); + } + } + else { + if (logger.isDebugEnabled()){ + logger.debug("Releasing " + item + " back to the pool"); + } + if (item != null) { + available.add(item); + } + permits.release(); + } + } + + public synchronized void removeAllIdleItems() { + T item; + while ((item = this.available.poll()) != null) { + doRemoveItem(item); + } + } + + private void doRemoveItem(T item) { + this.allocated.remove(item); + this.callback.removedFromPool(item); + } + + /** + * User of the pool provide an implementation of this interface; called during + * various pool operations. + * + */ + public static interface PoolItemCallback { + + /** + * Called by the pool when a new instance is required to populate the pool. Only + * called if no idle non-stale instances are available. + * @return The item. + */ + T createForPool(); + + /** + * Called by the pool when an idle item is retrieved from the pool. Indicates + * whether that item is usable, or should be discarded. The pool takes no + * further action on a stale item, discards it, and attempts to find or create + * another item. + * @param item The item. + * @return true if the item should not be used. + */ + boolean isStale(T item); + + /** + * Called by the pool when an item is forcibly removed from the pool - for example + * when the pool size is reduced. The implementation should perform any cleanup + * necessary on the item, such as closing connections etc. + * @param item The item. + */ + void removedFromPool(T item); + } +} diff --git a/spring-integration-core/src/test/java/org/springframework/integration/util/SimplePoolTests.java b/spring-integration-core/src/test/java/org/springframework/integration/util/SimplePoolTests.java new file mode 100644 index 0000000000..0b2569de99 --- /dev/null +++ b/spring-integration-core/src/test/java/org/springframework/integration/util/SimplePoolTests.java @@ -0,0 +1,152 @@ +/* + * Copyright 2002-2012 the original author or authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.springframework.integration.util; + +import static org.junit.Assert.assertEquals; +import static org.junit.Assert.assertFalse; +import static org.junit.Assert.assertNotSame; +import static org.junit.Assert.assertSame; +import static org.junit.Assert.fail; + +import java.util.HashSet; +import java.util.Set; +import java.util.concurrent.atomic.AtomicBoolean; + +import org.junit.Test; +import org.springframework.integration.MessagingException; + +/** + * @author Gary Russell + * @since 2.2 + * + */ +public class SimplePoolTests { + + @Test + public void testReuseAndStale() { + final Set strings = new HashSet(); + final AtomicBoolean stale = new AtomicBoolean(); + SimplePool pool = stringPool(2, strings, stale); + String s1 = pool.getItem(); + String s2 = pool.getItem(); + assertNotSame(s1, s2); + pool.releaseItem(s1); + String s3 = pool.getItem(); + assertSame(s1, s3); + stale.set(true); + pool.releaseItem(s3); + s3 = pool.getItem(); + assertNotSame(s1, s3); + assertFalse(strings.remove(s1)); + } + + @Test + public void testOverCommitandResize() { + final Set strings = new HashSet(); + final AtomicBoolean stale = new AtomicBoolean(); + SimplePool pool = stringPool(2, strings, stale); + String s1 = pool.getItem(); + assertEquals(0, pool.getIdleCount()); + assertEquals(1, pool.getActiveCount()); + assertEquals(1, pool.getAllocatedCount()); + pool.releaseItem(s1); + assertEquals(1, pool.getIdleCount()); + assertEquals(0, pool.getActiveCount()); + assertEquals(1, pool.getAllocatedCount()); + s1 = pool.getItem(); + assertEquals(0, pool.getIdleCount()); + assertEquals(1, pool.getActiveCount()); + assertEquals(1, pool.getAllocatedCount()); + String s2 = pool.getItem(); + assertNotSame(s1, s2); + pool.setWaitTimeout(1); + assertEquals(0, pool.getIdleCount()); + assertEquals(2, pool.getActiveCount()); + assertEquals(2, pool.getAllocatedCount()); + try { + pool.getItem(); + fail("Expected exception"); + } catch (MessagingException e) {} + + // resize up + pool.setPoolSize(4); + + assertEquals(0, pool.getIdleCount()); + assertEquals(2, pool.getActiveCount()); + assertEquals(2, pool.getAllocatedCount()); + String s3 = pool.getItem(); + String s4 = pool.getItem(); + assertEquals(0, pool.getIdleCount()); + assertEquals(4, pool.getActiveCount()); + assertEquals(4, pool.getAllocatedCount()); + pool.releaseItem(s4); + assertEquals(1, pool.getIdleCount()); + assertEquals(3, pool.getActiveCount()); + assertEquals(4, pool.getAllocatedCount()); + + // resize down + pool.setPoolSize(2); + + assertEquals(0, pool.getIdleCount()); + assertEquals(3, pool.getActiveCount()); + assertEquals(3, pool.getPoolSize()); + assertEquals(3, pool.getAllocatedCount()); + pool.releaseItem(s3); + assertEquals(0, pool.getIdleCount()); + assertEquals(2, pool.getActiveCount()); + assertEquals(2, pool.getPoolSize()); + assertEquals(2, pool.getAllocatedCount()); + assertEquals(2, strings.size()); + pool.releaseItem(s2); + pool.releaseItem(s1); + assertEquals(2, pool.getIdleCount()); + assertEquals(0, pool.getActiveCount()); + assertEquals(2, pool.getPoolSize()); + assertEquals(2, strings.size()); + assertEquals(2, pool.getAllocatedCount()); + } + + @Test(expected=IllegalArgumentException.class) + public void testForeignObject() { + final Set strings = new HashSet(); + final AtomicBoolean stale = new AtomicBoolean(); + SimplePool pool = stringPool(2, strings, stale); + pool.getItem(); + pool.releaseItem("Hello, world!"); + } + + private SimplePool stringPool(int size, final Set strings, + final AtomicBoolean stale) { + SimplePool pool = new SimplePool(size, new SimplePool.PoolItemCallback() { + private int i; + public String createForPool() { + String string = new String("String" + i++); + strings.add(string); + return string; + } + public boolean isStale(String item) { + if (stale.get()) { + strings.remove(item); + } + return stale.get(); + } + public void removedFromPool(String item) { + strings.remove(item); + } + }); + return pool; + } +} diff --git a/spring-integration-file/src/main/java/org/springframework/integration/file/remote/session/CachingSessionFactory.java b/spring-integration-file/src/main/java/org/springframework/integration/file/remote/session/CachingSessionFactory.java index d5a1664c57..23a416a625 100644 --- a/spring-integration-file/src/main/java/org/springframework/integration/file/remote/session/CachingSessionFactory.java +++ b/spring-integration-file/src/main/java/org/springframework/integration/file/remote/session/CachingSessionFactory.java @@ -19,12 +19,11 @@ package org.springframework.integration.file.remote.session; import java.io.IOException; import java.io.InputStream; import java.io.OutputStream; -import java.util.concurrent.LinkedBlockingQueue; import org.apache.commons.logging.Log; import org.apache.commons.logging.LogFactory; import org.springframework.beans.factory.DisposableBean; -import org.springframework.integration.util.UpperBound; +import org.springframework.integration.util.SimplePool; /** * A {@link SessionFactory} implementation that caches Sessions for reuse without @@ -41,15 +40,9 @@ public class CachingSessionFactory implements SessionFactory, DisposableBe private static final Log logger = LogFactory.getLog(CachingSessionFactory.class); - - private volatile long sessionWaitTimeout = Integer.MAX_VALUE; - - private final LinkedBlockingQueue> queue = new LinkedBlockingQueue>(); - private final SessionFactory sessionFactory; - private final UpperBound sessionPermits; - + private final SimplePool> pool; public CachingSessionFactory(SessionFactory sessionFactory) { this(sessionFactory, 0); @@ -57,7 +50,19 @@ public class CachingSessionFactory implements SessionFactory, DisposableBe public CachingSessionFactory(SessionFactory sessionFactory, int sessionCacheSize) { this.sessionFactory = sessionFactory; - this.sessionPermits = new UpperBound(sessionCacheSize); + this.pool = new SimplePool>(sessionCacheSize, new SimplePool.PoolItemCallback>() { + public Session createForPool() { + return CachingSessionFactory.this.sessionFactory.getSession(); + } + + public boolean isStale(Session session) { + return !session.isOpen(); + } + + public void removedFromPool(Session session) { + session.close(); + } + }); } @@ -67,52 +72,19 @@ public class CachingSessionFactory implements SessionFactory, DisposableBe * @throws {@link IllegalStateException} if the wait expires prior to a Session becoming available. */ public void setSessionWaitTimeout(long sessionWaitTimeout) { - this.sessionWaitTimeout = sessionWaitTimeout; + this.pool.setWaitTimeout(sessionWaitTimeout); } - public Session getSession() { - boolean permitted = this.sessionPermits.tryAcquire(this.sessionWaitTimeout); - if (!permitted) { - throw new IllegalStateException("Timed out while waiting to aquire a Session."); - } - Session session = this.doGetSession(); - return new CachedSession(session); + public void setPoolSize(int poolSize) { + this.pool.setPoolSize(poolSize); + } + + public Session getSession() { + return new CachedSession(this.pool.getItem()); } public void destroy() { - if (this.queue != null) { - for (Session session : this.queue) { - this.closeSession(session); - } - } - } - - private Session doGetSession() { - Session session = this.queue.poll(); - if (session != null && !session.isOpen()) { - if (logger.isDebugEnabled()) { - logger.debug("Received a stale Session, will attempt to get a new one."); - } - return this.doGetSession(); - } - else if (session == null){ - session = this.sessionFactory.getSession(); - } - return session; - } - - private void closeSession(Session session) { - try { - if (session != null) { - session.close(); - } - } - catch (Throwable e) { - if (logger.isWarnEnabled()) { - // log and ignore - logger.warn("Exception was thrown while destroying Session. ", e); - } - } + this.pool.removeAllIdleItems(); } @@ -120,16 +92,25 @@ public class CachingSessionFactory implements SessionFactory, DisposableBe private final Session targetSession; + private boolean released; + private CachedSession(Session targetSession) { this.targetSession = targetSession; } - public void close() { - if (logger.isDebugEnabled()){ - logger.debug("Releasing Session back to the pool."); + public synchronized void close() { + if (released) { + if (logger.isDebugEnabled()){ + logger.debug("Session already released."); + } + } + else { + if (logger.isDebugEnabled()){ + logger.debug("Releasing Session back to the pool."); + } + pool.releaseItem(targetSession); + released = true; } - queue.add(targetSession); - sessionPermits.release(); } public boolean remove(String path) throws IOException{ diff --git a/spring-integration-ftp/src/test/java/org/springframework/integration/ftp/session/SessionFactoryTests.java b/spring-integration-ftp/src/test/java/org/springframework/integration/ftp/session/SessionFactoryTests.java index fbc990ec53..24e3ef7377 100644 --- a/spring-integration-ftp/src/test/java/org/springframework/integration/ftp/session/SessionFactoryTests.java +++ b/spring-integration-ftp/src/test/java/org/springframework/integration/ftp/session/SessionFactoryTests.java @@ -27,6 +27,7 @@ import org.apache.commons.net.ftp.FTPClient; import org.junit.Ignore; import org.junit.Test; import org.mockito.Mockito; +import org.springframework.integration.MessagingException; import org.springframework.integration.file.remote.session.CachingSessionFactory; import org.springframework.integration.file.remote.session.Session; import org.springframework.integration.file.remote.session.SessionFactory; @@ -133,7 +134,7 @@ public class SessionFactoryTests { Mockito.verify(sessionFactory, Mockito.times(2)).getSession(); } - @Test (expected=IllegalStateException.class) // timeout expire + @Test (expected=MessagingException.class) // timeout expire public void testSessionWaitExpire() throws Exception{ SessionFactory sessionFactory = Mockito.mock(SessionFactory.class); Session session = Mockito.mock(Session.class); diff --git a/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/connection/CachingClientConnectionFactory.java b/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/connection/CachingClientConnectionFactory.java new file mode 100644 index 0000000000..4f08b4494d --- /dev/null +++ b/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/connection/CachingClientConnectionFactory.java @@ -0,0 +1,323 @@ +/* + * Copyright 2002-2012 the original author or authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.springframework.integration.ip.tcp.connection; + +import java.util.concurrent.Executor; + +import org.apache.commons.logging.Log; +import org.apache.commons.logging.LogFactory; +import org.springframework.core.serializer.Deserializer; +import org.springframework.core.serializer.Serializer; +import org.springframework.integration.MessagingException; +import org.springframework.integration.util.SimplePool; + +/** + * @author Gary Russell + * @since 2.2 + * + */ +public class CachingClientConnectionFactory extends AbstractClientConnectionFactory { + + private final Log logger = LogFactory.getLog(this.getClass()); + + private final AbstractClientConnectionFactory targetConnectionFactory; + + private final SimplePool pool; + + public CachingClientConnectionFactory(AbstractClientConnectionFactory target, int poolSize) { + super("", 0); + // override single-use to true to force "close" after use + target.setSingleUse(true); + this.targetConnectionFactory = target; + pool = new SimplePool(poolSize, new SimplePool.PoolItemCallback() { + + public TcpConnection createForPool() { + try { + return targetConnectionFactory.getConnection(); + } catch (Exception e) { + throw new MessagingException("Failed to obtain connection", e); + } + } + + public boolean isStale(TcpConnection connection) { + return !connection.isOpen(); + } + + public void removedFromPool(TcpConnection connection) { + connection.close(); + } + }); + } + + public void setConnectionWaitTimeout(int connectionWaitTimeout) { + this.pool.setWaitTimeout(connectionWaitTimeout); + } + + public synchronized void setPoolSize(int poolSize) { + this.pool.setPoolSize(poolSize); + } + + public int getPoolSize() { + return this.pool.getPoolSize(); + } + + public int getIdleCount() { + return this.pool.getIdleCount(); + } + + public int getActiveCount() { + return this.pool.getActiveCount(); + } + + public int getAllocatedCount() { + return this.pool.getAllocatedCount(); + } + + public TcpConnection getOrMakeConnection() throws Exception { + return new CachedConnection(this.pool.getItem()); + } + + private class CachedConnection extends AbstractTcpConnectionInterceptor { + + private volatile boolean released; + + public CachedConnection(TcpConnection connection) { + super.setTheConnection(connection); + if (connection instanceof AbstractTcpConnection) { + ((AbstractTcpConnection) connection).registerListener(this); + } + } + + @Override + public synchronized void close() { + /** + * If the delegate is stopped, actually close + * the connection. + */ + if (!isRunning()) { + if (logger.isDebugEnabled()){ + logger.debug("Factory not running - closing " + this.getConnectionId()); + } + pool.releaseItem(null); // just open up a permit + super.close(); + } + else if(this.released) { + if (logger.isDebugEnabled()) { + logger.debug("Connection " + this.getConnectionId() + " has already been released"); + } + } + else { + pool.releaseItem(this.getTheConnection()); + this.released = true; + } + } + + @Override + public String getConnectionId() { + return "Cached:" + super.getConnectionId(); + } + + @Override + public String toString() { + return this.getConnectionId(); + } + + } + +///////////////// DELEGATE METHODS /////////////////////// + + public void run() { + + } + + public boolean isRunning() { + return targetConnectionFactory.isRunning(); + } + + @Override + public void close() { + targetConnectionFactory.close(); + } + + public int hashCode() { + return targetConnectionFactory.hashCode(); + } + + public void setComponentName(String componentName) { + targetConnectionFactory.setComponentName(componentName); + } + + public String getComponentType() { + return targetConnectionFactory.getComponentType(); + } + + public boolean equals(Object obj) { + return targetConnectionFactory.equals(obj); + } + + public int getSoTimeout() { + return targetConnectionFactory.getSoTimeout(); + } + + public void setSoTimeout(int soTimeout) { + targetConnectionFactory.setSoTimeout(soTimeout); + } + + public int getSoReceiveBufferSize() { + return targetConnectionFactory.getSoReceiveBufferSize(); + } + + public void setSoReceiveBufferSize(int soReceiveBufferSize) { + targetConnectionFactory.setSoReceiveBufferSize(soReceiveBufferSize); + } + + public int getSoSendBufferSize() { + return targetConnectionFactory.getSoSendBufferSize(); + } + + public void setSoSendBufferSize(int soSendBufferSize) { + targetConnectionFactory.setSoSendBufferSize(soSendBufferSize); + } + + public boolean isSoTcpNoDelay() { + return targetConnectionFactory.isSoTcpNoDelay(); + } + + public void setSoTcpNoDelay(boolean soTcpNoDelay) { + targetConnectionFactory.setSoTcpNoDelay(soTcpNoDelay); + } + + public int getSoLinger() { + return targetConnectionFactory.getSoLinger(); + } + + public void setSoLinger(int soLinger) { + targetConnectionFactory.setSoLinger(soLinger); + } + + public boolean isSoKeepAlive() { + return targetConnectionFactory.isSoKeepAlive(); + } + + public void setSoKeepAlive(boolean soKeepAlive) { + targetConnectionFactory.setSoKeepAlive(soKeepAlive); + } + + public int getSoTrafficClass() { + return targetConnectionFactory.getSoTrafficClass(); + } + + public void setSoTrafficClass(int soTrafficClass) { + targetConnectionFactory.setSoTrafficClass(soTrafficClass); + } + + public String getHost() { + return targetConnectionFactory.getHost(); + } + + public int getPort() { + return targetConnectionFactory.getPort(); + } + + public TcpListener getListener() { + return targetConnectionFactory.getListener(); + } + + public TcpSender getSender() { + return targetConnectionFactory.getSender(); + } + + public Serializer getSerializer() { + return targetConnectionFactory.getSerializer(); + } + + public Deserializer getDeserializer() { + return targetConnectionFactory.getDeserializer(); + } + + public TcpMessageMapper getMapper() { + return targetConnectionFactory.getMapper(); + } + + public void registerListener(TcpListener listener) { + targetConnectionFactory.registerListener(listener); + } + + public void registerSender(TcpSender sender) { + targetConnectionFactory.registerSender(sender); + } + + public void setTaskExecutor(Executor taskExecutor) { + targetConnectionFactory.setTaskExecutor(taskExecutor); + } + + public void setDeserializer(Deserializer deserializer) { + targetConnectionFactory.setDeserializer(deserializer); + } + + public void setSerializer(Serializer serializer) { + targetConnectionFactory.setSerializer(serializer); + } + + public void setMapper(TcpMessageMapper mapper) { + targetConnectionFactory.setMapper(mapper); + } + + public boolean isSingleUse() { + return targetConnectionFactory.isSingleUse(); + } + + public void setSingleUse(boolean singleUse) { + targetConnectionFactory.setSingleUse(singleUse); + } + + public void setInterceptorFactoryChain( + TcpConnectionInterceptorFactoryChain interceptorFactoryChain) { + targetConnectionFactory + .setInterceptorFactoryChain(interceptorFactoryChain); + } + + public void setLookupHost(boolean lookupHost) { + targetConnectionFactory.setLookupHost(lookupHost); + } + + public boolean isLookupHost() { + return targetConnectionFactory.isLookupHost(); + } + + public void start() { + this.setActive(true); + targetConnectionFactory.start(); + } + + public synchronized void stop() { + targetConnectionFactory.stop(); + this.pool.removeAllIdleItems(); + } + + public int getPhase() { + return targetConnectionFactory.getPhase(); + } + + public boolean isAutoStartup() { + return targetConnectionFactory.isAutoStartup(); + } + + public void stop(Runnable callback) { + targetConnectionFactory.stop(callback); + } + +} diff --git a/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/connection/CachingClientConnectionFactoryTests-context.xml b/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/connection/CachingClientConnectionFactoryTests-context.xml new file mode 100644 index 0000000000..33210baecd --- /dev/null +++ b/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/connection/CachingClientConnectionFactoryTests-context.xml @@ -0,0 +1,46 @@ + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/connection/CachingClientConnectionFactoryTests.java b/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/connection/CachingClientConnectionFactoryTests.java new file mode 100644 index 0000000000..719e945da2 --- /dev/null +++ b/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/connection/CachingClientConnectionFactoryTests.java @@ -0,0 +1,293 @@ +/* + * Copyright 2002-2012 the original author or authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.springframework.integration.ip.tcp.connection; + +import static org.junit.Assert.assertEquals; +import static org.junit.Assert.assertNotNull; +import static org.junit.Assert.assertNotSame; +import static org.junit.Assert.assertSame; +import static org.junit.Assert.fail; +import static org.mockito.Mockito.doAnswer; +import static org.mockito.Mockito.doThrow; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.times; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import java.util.concurrent.Semaphore; + +import org.junit.Test; +import org.junit.runner.RunWith; +import org.mockito.invocation.InvocationOnMock; +import org.mockito.stubbing.Answer; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.integration.Message; +import org.springframework.integration.MessagingException; +import org.springframework.integration.core.PollableChannel; +import org.springframework.integration.core.SubscribableChannel; +import org.springframework.integration.ip.IpHeaders; +import org.springframework.integration.message.GenericMessage; +import org.springframework.integration.test.util.TestUtils; +import org.springframework.test.annotation.ExpectedException; +import org.springframework.test.context.ContextConfiguration; +import org.springframework.test.context.junit4.SpringJUnit4ClassRunner; + +/** + * @author Gary Russell + * @since 2.2 + * + */ +@ContextConfiguration +@RunWith(SpringJUnit4ClassRunner.class) +public class CachingClientConnectionFactoryTests { + + @Autowired + SubscribableChannel outbound; + + @Autowired + PollableChannel inbound; + + @Autowired + AbstractServerConnectionFactory serverCf; + + @Test + public void testReuse() throws Exception { + AbstractClientConnectionFactory factory = mock(AbstractClientConnectionFactory.class); + when(factory.isRunning()).thenReturn(true); + TcpConnection mockConn1 = makeMockConnection("conn1"); + TcpConnection mockConn2 = makeMockConnection("conn2"); + when(factory.getConnection()).thenReturn(mockConn1).thenReturn(mockConn2); + CachingClientConnectionFactory cachingFactory = new CachingClientConnectionFactory(factory, 2); + cachingFactory.start(); + TcpConnection conn1 = cachingFactory.getConnection(); + assertEquals("Cached:" + mockConn1.toString(), conn1.toString()); + conn1.close(); + conn1 = cachingFactory.getConnection(); + assertEquals("Cached:" + mockConn1.toString(), conn1.toString()); + TcpConnection conn2 = cachingFactory.getConnection(); + assertEquals("Cached:" + mockConn2.toString(), conn2.toString()); + conn1.close(); + conn2.close(); + } + + @Test + public void testReuseNoLimit() throws Exception { + AbstractClientConnectionFactory factory = mock(AbstractClientConnectionFactory.class); + when(factory.isRunning()).thenReturn(true); + TcpConnection mockConn1 = makeMockConnection("conn1"); + TcpConnection mockConn2 = makeMockConnection("conn2"); + when(factory.getConnection()).thenReturn(mockConn1).thenReturn(mockConn2); + CachingClientConnectionFactory cachingFactory = new CachingClientConnectionFactory(factory, 0); + cachingFactory.start(); + TcpConnection conn1 = cachingFactory.getConnection(); + assertEquals("Cached:" + mockConn1.toString(), conn1.toString()); + conn1.close(); + conn1 = cachingFactory.getConnection(); + assertEquals("Cached:" + mockConn1.toString(), conn1.toString()); + TcpConnection conn2 = cachingFactory.getConnection(); + assertEquals("Cached:" + mockConn2.toString(), conn2.toString()); + conn1.close(); + conn2.close(); + } + + @Test + public void testReuseClosed() throws Exception { + AbstractClientConnectionFactory factory = mock(AbstractClientConnectionFactory.class); + when(factory.isRunning()).thenReturn(true); + TcpConnection mockConn1 = makeMockConnection("conn1"); + TcpConnection mockConn2 = makeMockConnection("conn2"); + when(factory.getConnection()).thenReturn(mockConn1) + .thenReturn(mockConn2).thenReturn(mockConn1) + .thenReturn(mockConn2); + CachingClientConnectionFactory cachingFactory = new CachingClientConnectionFactory(factory, 2); + cachingFactory.start(); + TcpConnection conn1 = cachingFactory.getConnection(); + assertEquals("Cached:" + mockConn1.toString(), conn1.toString()); + conn1.close(); + conn1 = cachingFactory.getConnection(); + assertEquals("Cached:" + mockConn1.toString(), conn1.toString()); + TcpConnection conn2 = cachingFactory.getConnection(); + assertEquals("Cached:" + mockConn2.toString(), conn2.toString()); + conn1.close(); + conn2.close(); + when(mockConn1.isOpen()).thenReturn(false); + TcpConnection conn2a = cachingFactory.getConnection(); + assertEquals("Cached:" + mockConn2.toString(), conn2a.toString()); + assertSame(TestUtils.getPropertyValue(conn2, "theConnection"), + TestUtils.getPropertyValue(conn2a, "theConnection")); + conn2a.close(); + } + + @Test @ExpectedException(MessagingException.class) + public void testLimit() throws Exception { + AbstractClientConnectionFactory factory = mock(AbstractClientConnectionFactory.class); + when(factory.isRunning()).thenReturn(true); + TcpConnection mockConn1 = makeMockConnection("conn1"); + TcpConnection mockConn2 = makeMockConnection("conn2"); + when(factory.getConnection()).thenReturn(mockConn1).thenReturn(mockConn2); + CachingClientConnectionFactory cachingFactory = new CachingClientConnectionFactory(factory, 2); + cachingFactory.setConnectionWaitTimeout(10); + cachingFactory.start(); + TcpConnection conn1 = cachingFactory.getConnection(); + assertEquals("Cached:" + mockConn1.toString(), conn1.toString()); + conn1.close(); + conn1 = cachingFactory.getConnection(); + assertEquals("Cached:" + mockConn1.toString(), conn1.toString()); + TcpConnection conn2 = cachingFactory.getConnection(); + assertEquals("Cached:" + mockConn2.toString(), conn2.toString()); + cachingFactory.getConnection(); + } + + @Test + public void testStop() throws Exception { + AbstractClientConnectionFactory factory = mock(AbstractClientConnectionFactory.class); + when(factory.isRunning()).thenReturn(true); + TcpConnection mockConn1 = makeMockConnection("conn1"); + TcpConnection mockConn2 = makeMockConnection("conn2"); + int i = 3; + when(factory.getConnection()).thenReturn(mockConn1) + .thenReturn(mockConn2) + .thenReturn(makeMockConnection("conn" + (i++))); + CachingClientConnectionFactory cachingFactory = new CachingClientConnectionFactory(factory, 2); + cachingFactory.start(); + TcpConnection conn1 = cachingFactory.getConnection(); + assertEquals("Cached:" + mockConn1.toString(), conn1.toString()); + conn1.close(); + conn1 = cachingFactory.getConnection(); + assertEquals("Cached:" + mockConn1.toString(), conn1.toString()); + TcpConnection conn2 = cachingFactory.getConnection(); + assertEquals("Cached:" + mockConn2.toString(), conn2.toString()); + cachingFactory.stop(); + Answer answer = new Answer () { + public Object answer(InvocationOnMock invocation) throws Throwable { + return null; + }}; + doAnswer(answer).when(mockConn1).close(); + doAnswer(answer).when(mockConn2).close(); + when(factory.isRunning()).thenReturn(false); + conn1.close(); + conn2.close(); + verify(mockConn1).close(); + verify(mockConn2).close(); + when(factory.isRunning()).thenReturn(true); + TcpConnection conn3 = cachingFactory.getConnection(); + assertNotSame(TestUtils.getPropertyValue(conn1, "theConnection"), + TestUtils.getPropertyValue(conn3, "theConnection")); + assertNotSame(TestUtils.getPropertyValue(conn2, "theConnection"), + TestUtils.getPropertyValue(conn3, "theConnection")); + } + + @Test + public void testEnlargePool() throws Exception { + AbstractClientConnectionFactory factory = mock(AbstractClientConnectionFactory.class); + when(factory.isRunning()).thenReturn(true); + TcpConnection mockConn = makeMockConnection("conn"); + when(factory.getConnection()).thenReturn(mockConn); + CachingClientConnectionFactory cachingFactory = new CachingClientConnectionFactory(factory, 2); + cachingFactory.start(); + TcpConnection conn1 = cachingFactory.getConnection(); + TcpConnection conn2 = cachingFactory.getConnection(); + assertNotSame(conn1, conn2); + Semaphore semaphore = TestUtils.getPropertyValue( + TestUtils.getPropertyValue(cachingFactory, "pool"), "permits", Semaphore.class); + assertEquals(0, semaphore.availablePermits()); + cachingFactory.setPoolSize(4); + TcpConnection conn3 = cachingFactory.getConnection(); + TcpConnection conn4 = cachingFactory.getConnection(); + assertEquals(0, semaphore.availablePermits()); + conn1.close(); + conn1.close(); + conn2.close(); + conn3.close(); + conn4.close(); + assertEquals(4, semaphore.availablePermits()); + } + + @Test + public void testReducePool() throws Exception { + AbstractClientConnectionFactory factory = mock(AbstractClientConnectionFactory.class); + when(factory.isRunning()).thenReturn(true); + TcpConnection mockConn1 = makeMockConnection("conn", true); + TcpConnection mockConn2 = makeMockConnection("conn", true); + TcpConnection mockConn3 = makeMockConnection("conn", true); + TcpConnection mockConn4 = makeMockConnection("conn", true); + when(factory.getConnection()).thenReturn(mockConn1) + .thenReturn(mockConn2).thenReturn(mockConn3) + .thenReturn(mockConn4); + CachingClientConnectionFactory cachingFactory = new CachingClientConnectionFactory(factory, 4); + cachingFactory.start(); + TcpConnection conn1 = cachingFactory.getConnection(); + TcpConnection conn2 = cachingFactory.getConnection(); + TcpConnection conn3 = cachingFactory.getConnection(); + TcpConnection conn4 = cachingFactory.getConnection(); + Semaphore semaphore = TestUtils.getPropertyValue( + TestUtils.getPropertyValue(cachingFactory, "pool"), "permits", Semaphore.class); + assertEquals(0, semaphore.availablePermits()); + conn1.close(); + assertEquals(1, semaphore.availablePermits()); + cachingFactory.setPoolSize(2); + assertEquals(0, semaphore.availablePermits()); + assertEquals(3, cachingFactory.getActiveCount()); + conn2.close(); + assertEquals(0, semaphore.availablePermits()); + assertEquals(2, cachingFactory.getActiveCount()); + conn3.close(); + assertEquals(1, cachingFactory.getActiveCount()); + assertEquals(1, cachingFactory.getIdleCount()); + conn4.close(); + assertEquals(2, semaphore.availablePermits()); + assertEquals(0, cachingFactory.getActiveCount()); + assertEquals(2, cachingFactory.getIdleCount()); + verify(mockConn1).close(); + verify(mockConn2).close(); + } + + private TcpConnection makeMockConnection(String name) { + return makeMockConnection(name, false); + } + + private TcpConnection makeMockConnection(String name, boolean closeOk) { + TcpConnection mockConn1 = mock(TcpConnection.class); + when(mockConn1.getConnectionId()).thenReturn(name); + when(mockConn1.toString()).thenReturn(name); + when(mockConn1.isOpen()).thenReturn(true); + if (!closeOk) { + doThrow(new RuntimeException("close() not expected")).when(mockConn1).close(); + } + return mockConn1; + } + + @Test + public void integrationTest() throws Exception { + int n = 0; + while (!serverCf.isListening()) { + Thread.sleep(100); + n++; + if (n > 10000) { + fail("Server didn't begin listening"); + } + } + outbound.send(new GenericMessage("Hello, world!")); + Message m = inbound.receive(1000); + assertNotNull(m); + String connectionId = m.getHeaders().get(IpHeaders.CONNECTION_ID, String.class); + outbound.send(new GenericMessage("Hello, world!")); + m = inbound.receive(1000); + assertNotNull(m); + assertEquals(connectionId, m.getHeaders().get(IpHeaders.CONNECTION_ID, String.class)); + + } +}