Fix CachedSessionFactory Race

Close the pool so that any sessions returned after the factory is
`destroy()`ed are closed.

* Call `removeAllIdleItems()` in `close()`.

* Close sessions in `SftpStreamingMessageSourceTests`.

**cherry-pick to all supported branches**
# Conflicts:
#	spring-integration-core/src/test/java/org/springframework/integration/util/SimplePoolTests.java
This commit is contained in:
Gary Russell
2020-07-08 14:19:57 -04:00
committed by Artem Bilan
parent fe524dd1e6
commit 8df8028537
6 changed files with 63 additions and 6 deletions

View File

@@ -72,4 +72,11 @@ public interface Pool<T> {
*/
int getAllocatedCount();
/**
* Close the pool; returned items will be destroyed.
* @since 4.3.23
*/
default void close() {
}
}

View File

@@ -59,6 +59,8 @@ public class SimplePool<T> implements Pool<T> {
private final PoolItemCallback<T> callback;
private volatile boolean closed;
/**
* Creates a SimplePool with a specific limit.
* @param poolSize The maximum number of items the pool supports.
@@ -158,6 +160,7 @@ public class SimplePool<T> implements Pool<T> {
*/
@Override
public T getItem() {
Assert.state(!this.closed, "Pool has been closed");
boolean permitted = false;
try {
try {
@@ -215,7 +218,7 @@ public class SimplePool<T> implements Pool<T> {
Assert.isTrue(this.allocated.contains(item),
"You can only release items that were obtained from the pool");
if (this.inUse.contains(item)) {
if (this.poolSize.get() > this.targetPoolSize.get()) {
if (this.poolSize.get() > this.targetPoolSize.get() || this.closed) {
this.poolSize.decrementAndGet();
if (item != null) {
doRemoveItem(item);
@@ -256,6 +259,12 @@ public class SimplePool<T> implements Pool<T> {
this.callback.removedFromPool(item);
}
@Override
public synchronized void close() {
this.closed = true;
removeAllIdleItems();
}
/**
* User of the pool provide an implementation of this interface; called during
* various pool operations.

View File

@@ -1,5 +1,5 @@
/*
* Copyright 2002-2019 the original author or authors.
* Copyright 2002-2020 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.
@@ -16,10 +16,12 @@
package org.springframework.integration.util;
import static org.hamcrest.Matchers.instanceOf;
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.assertThat;
import static org.junit.Assert.fail;
import java.util.HashSet;
@@ -27,6 +29,7 @@ import java.util.Set;
import java.util.concurrent.Semaphore;
import java.util.concurrent.atomic.AtomicBoolean;
import org.junit.Test;
import org.springframework.integration.test.util.TestUtils;
@@ -124,13 +127,18 @@ public class SimplePoolTests {
assertEquals(2, pool.getAllocatedCount());
}
@Test(expected = IllegalArgumentException.class)
@Test
public void testForeignObject() {
final Set<String> strings = new HashSet<String>();
final AtomicBoolean stale = new AtomicBoolean();
SimplePool<String> pool = stringPool(2, strings, stale);
pool.getItem();
pool.releaseItem("Hello, world!");
try {
pool.releaseItem("Hello, world!");
}
catch (Exception e) {
assertThat(e, instanceOf(IllegalArgumentException.class));
}
}
@Test
@@ -149,16 +157,37 @@ public class SimplePoolTests {
}
@Test
public void testClose() {
SimplePool<String> pool = stringPool(10, new HashSet<>(), new AtomicBoolean());
String item1 = pool.getItem();
String item2 = pool.getItem();
pool.releaseItem(item2);
assertEquals(2, pool.getAllocatedCount());
pool.close();
pool.releaseItem(item1);
assertEquals(0, pool.getAllocatedCount());
try {
pool.getItem();
}
catch (Exception e) {
assertThat(e, instanceOf(IllegalStateException.class));
}
}
private SimplePool<String> stringPool(int size, final Set<String> strings,
final AtomicBoolean stale) {
SimplePool<String> pool = new SimplePool<String>(size, new SimplePool.PoolItemCallback<String>() {
private int i;
@Override
public String createForPool() {
String string = "String" + i++;
strings.add(string);
return string;
}
@Override
public boolean isStale(String item) {
if (stale.get()) {
@@ -166,10 +195,12 @@ public class SimplePoolTests {
}
return stale.get();
}
@Override
public void removedFromPool(String item) {
strings.remove(item);
}
});
return pool;
}

View File

@@ -140,7 +140,7 @@ public class CachingSessionFactory<F> implements SessionFactory<F>, DisposableBe
*/
@Override
public void destroy() {
this.pool.removeAllIdleItems();
this.pool.close();
}
/**

View File

@@ -21,6 +21,7 @@ import java.util.Map;
import java.util.concurrent.Executor;
import java.util.concurrent.atomic.AtomicBoolean;
import org.springframework.beans.factory.DisposableBean;
import org.springframework.core.serializer.Deserializer;
import org.springframework.core.serializer.Serializer;
import org.springframework.integration.ip.IpHeaders;
@@ -41,7 +42,7 @@ import org.springframework.messaging.support.ErrorMessage;
* @since 2.2
*
*/
public class CachingClientConnectionFactory extends AbstractClientConnectionFactory {
public class CachingClientConnectionFactory extends AbstractClientConnectionFactory implements DisposableBean {
private final AbstractClientConnectionFactory targetConnectionFactory;
@@ -389,6 +390,11 @@ public class CachingClientConnectionFactory extends AbstractClientConnectionFact
this.pool.removeAllIdleItems();
}
@Override
public void destroy() throws Exception {
this.pool.close();
}
private final class CachedConnection extends TcpConnectionInterceptorSupport {
private final AtomicBoolean released = new AtomicBoolean();

View File

@@ -36,6 +36,7 @@ import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.context.ApplicationContext;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.springframework.integration.StaticMessageHeaderAccessor;
import org.springframework.integration.annotation.InboundChannelAdapter;
import org.springframework.integration.annotation.Transformer;
import org.springframework.integration.channel.QueueChannel;
@@ -138,6 +139,7 @@ public class SftpStreamingMessageSourceTests extends SftpTestSupport {
anyOf(equalTo(" sftpSource1.txt"), equalTo("sftpSource2.txt")));
received.getPayload().close();
StaticMessageHeaderAccessor.getCloseableResource(received).close();
}
@Test
@@ -152,6 +154,7 @@ public class SftpStreamingMessageSourceTests extends SftpTestSupport {
anyOf(equalTo(" sftpSource1.txt"), equalTo("sftpSource2.txt")));
received.getPayload().close();
StaticMessageHeaderAccessor.getCloseableResource(received).close();
}
@Test
@@ -166,6 +169,7 @@ public class SftpStreamingMessageSourceTests extends SftpTestSupport {
anyOf(equalTo(" sftpSource1.txt"), equalTo("sftpSource2.txt")));
received.getPayload().close();
StaticMessageHeaderAccessor.getCloseableResource(received).close();
}
private SftpStreamingMessageSource buildSource() {