From 0f46176c3d2bd2c4a713054f6c27da973386e346 Mon Sep 17 00:00:00 2001 From: Gary Russell Date: Tue, 29 Nov 2011 19:23:59 -0500 Subject: [PATCH] INT-2287 Timing Hole With Send-and-Forget TCP If a tcp client connection factory was configured for single-use connections (one socket per message), and it was used by an outbound channel adapter (send and forget), the socket is closed immediately after sending the message. This could cause issues with the connection handling because certain failures could occur. For example, the close could occur before we ever registered the channel for read selection, resulting in a ClosedChannelException. Secondly, there was a small timing hole where a selection key could be invalidated between the isValid() call and isReadable(). These problems have been corrected. --- .../connection/AbstractConnectionFactory.java | 38 ++++++++-------- .../TcpNioClientConnectionFactory.java | 9 +++- .../ip/tcp/connection/TcpNioConnection.java | 3 ++ .../ip/tcp/TcpSendingMessageHandlerTests.java | 45 ++++++++++++------- 4 files changed, 58 insertions(+), 37 deletions(-) diff --git a/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/connection/AbstractConnectionFactory.java b/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/connection/AbstractConnectionFactory.java index 480e1bb060..51d8ff8fc7 100644 --- a/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/connection/AbstractConnectionFactory.java +++ b/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/connection/AbstractConnectionFactory.java @@ -532,11 +532,11 @@ public abstract class AbstractConnectionFactory extends IntegrationObjectSupport while (iterator.hasNext()) { final SelectionKey key = iterator.next(); iterator.remove(); - if (!key.isValid()) { - logger.debug("Selection key no longer valid"); - } - else if (key.isReadable()) { - try { + try { + if (!key.isValid()) { + logger.debug("Selection key no longer valid"); + } + else if (key.isReadable()) { key.interestOps(key.interestOps() - key.readyOps()); final TcpNioConnection connection; connection = (TcpNioConnection) key.attachment(); @@ -560,23 +560,23 @@ public abstract class AbstractConnectionFactory extends IntegrationObjectSupport selector.wakeup(); } }}); - } catch (Exception e) { - if (e instanceof CancelledKeyException) { - logger.debug("Exception on readable key", e); - continue; + } + else if (key.isAcceptable()) { + try { + doAccept(selector, server, now); + } catch (Exception e) { + logger.error("Exception accepting new connection", e); } - logger.error("Exception on readable key", e); } - } - else if (key.isAcceptable()) { - try { - doAccept(selector, server, now); - } catch (Exception e) { - logger.error("Exception accepting new connection", e); + else { + logger.error("Unexpected key: " + key); } - } - else { - logger.error("Unexpected key: " + key); + } catch (CancelledKeyException e) { + if (logger.isDebugEnabled()) { + logger.debug("Selection key " + key + " cancelled"); + } + } catch (Exception e) { + logger.error("Exception on selection key " + key, e); } } } diff --git a/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/connection/TcpNioClientConnectionFactory.java b/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/connection/TcpNioClientConnectionFactory.java index 37aba71098..5de0da2c9c 100644 --- a/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/connection/TcpNioClientConnectionFactory.java +++ b/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/connection/TcpNioClientConnectionFactory.java @@ -20,6 +20,7 @@ import java.io.IOException; import java.net.InetSocketAddress; import java.net.SocketException; import java.nio.ByteBuffer; +import java.nio.channels.ClosedChannelException; import java.nio.channels.SelectionKey; import java.nio.channels.Selector; import java.nio.channels.SocketChannel; @@ -124,7 +125,13 @@ public class TcpNioClientConnectionFactory extends int soTimeout = this.getSoTimeout(); int selectionCount = selector.select(soTimeout < 0 ? 0 : soTimeout); while ((newChannel = newChannels.poll()) != null) { - newChannel.register(this.selector, SelectionKey.OP_READ, connections.get(newChannel)); + try { + newChannel.register(this.selector, SelectionKey.OP_READ, connections.get(newChannel)); + } catch (ClosedChannelException cce) { + if (logger.isDebugEnabled()) { + logger.debug("Channel closed before registering with selector for reading"); + } + } } this.processNioSelections(selectionCount, selector, null, this.connections); } diff --git a/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/connection/TcpNioConnection.java b/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/connection/TcpNioConnection.java index f3d4701cc5..fe1b94ac6e 100644 --- a/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/connection/TcpNioConnection.java +++ b/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/connection/TcpNioConnection.java @@ -306,6 +306,9 @@ public class TcpNioConnection extends AbstractTcpConnection { try { doRead(); } catch (ClosedChannelException cce) { + if (logger.isDebugEnabled()) { + logger.debug(this.getConnectionId() + " Channel is closed"); + } this.closeConnection(); } catch (Exception e) { logger.error("Exception on Read " + diff --git a/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/TcpSendingMessageHandlerTests.java b/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/TcpSendingMessageHandlerTests.java index 7a55e32e4e..aa0eb3e0e4 100644 --- a/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/TcpSendingMessageHandlerTests.java +++ b/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/TcpSendingMessageHandlerTests.java @@ -43,7 +43,6 @@ import javax.net.ServerSocketFactory; import org.apache.commons.logging.Log; import org.apache.commons.logging.LogFactory; import org.junit.Test; - import org.springframework.core.serializer.DefaultDeserializer; import org.springframework.core.serializer.DefaultSerializer; import org.springframework.integration.Message; @@ -552,13 +551,15 @@ public class TcpSendingMessageHandlerTests { try { ServerSocket server = ServerSocketFactory.getDefault().createServerSocket(port); latch.countDown(); - while (true) { + for (int i = 0; i < 2; i++) { Socket socket = server.accept(); semaphore.release(); byte[] b = new byte[6]; readFully(socket.getInputStream(), b); semaphore.release(); + socket.close(); } + server.close(); } catch (Exception e) { if (!done.get()) { e.printStackTrace(); @@ -580,6 +581,7 @@ public class TcpSendingMessageHandlerTests { handler.handleMessage(MessageBuilder.withPayload("Test").build()); assertTrue(semaphore.tryAcquire(4, 10000, TimeUnit.MILLISECONDS)); done.set(true); + ccf.stop(); } @Test @@ -593,13 +595,15 @@ public class TcpSendingMessageHandlerTests { try { ServerSocket server = ServerSocketFactory.getDefault().createServerSocket(port); latch.countDown(); - while (true) { + for (int i = 0; i < 2; i++) { Socket socket = server.accept(); semaphore.release(); - byte[] b = new byte[6]; + byte[] b = new byte[8]; readFully(socket.getInputStream(), b); semaphore.release(); + socket.close(); } + server.close(); } catch (Exception e) { if (!done.get()) { e.printStackTrace(); @@ -611,16 +615,17 @@ public class TcpSendingMessageHandlerTests { ByteArrayCrLfSerializer serializer = new ByteArrayCrLfSerializer(); ccf.setSerializer(serializer); ccf.setDeserializer(serializer); - ccf.setSoTimeout(10000); + ccf.setSoTimeout(5000); ccf.start(); ccf.setSingleUse(true); TcpSendingMessageHandler handler = new TcpSendingMessageHandler(); handler.setConnectionFactory(ccf); assertTrue(latch.await(10, TimeUnit.SECONDS)); - handler.handleMessage(MessageBuilder.withPayload("Test").build()); - handler.handleMessage(MessageBuilder.withPayload("Test").build()); + handler.handleMessage(MessageBuilder.withPayload("Test.1").build()); + handler.handleMessage(MessageBuilder.withPayload("Test.2").build()); assertTrue(semaphore.tryAcquire(4, 10000, TimeUnit.MILLISECONDS)); done.set(true); + ccf.stop(); } @Test @@ -634,15 +639,16 @@ public class TcpSendingMessageHandlerTests { try { ServerSocket server = ServerSocketFactory.getDefault().createServerSocket(port); latch.countDown(); - int i = 0; - while (true) { + for (int i = 1; i < 3; i++) { Socket socket = server.accept(); semaphore.release(); byte[] b = new byte[6]; readFully(socket.getInputStream(), b); - b = ("Reply" + (++i) + "\r\n").getBytes(); + b = ("Reply" + i + "\r\n").getBytes(); socket.getOutputStream().write(b); + socket.close(); } + server.close(); } catch (Exception e) { if (!done.get()) { e.printStackTrace(); @@ -676,6 +682,7 @@ public class TcpSendingMessageHandlerTests { assertTrue(replies.remove("Reply1")); assertTrue(replies.remove("Reply2")); done.set(true); + ccf.stop(); } @Test @@ -689,15 +696,16 @@ public class TcpSendingMessageHandlerTests { try { ServerSocket server = ServerSocketFactory.getDefault().createServerSocket(port); latch.countDown(); - int i = 0; - while (true) { + for (int i = 1; i < 3; i++) { Socket socket = server.accept(); semaphore.release(); byte[] b = new byte[6]; readFully(socket.getInputStream(), b); - b = ("Reply" + (++i) + "\r\n").getBytes(); + b = ("Reply" + i + "\r\n").getBytes(); socket.getOutputStream().write(b); + socket.close(); } + server.close(); } catch (Exception e) { if (!done.get()) { e.printStackTrace(); @@ -731,6 +739,7 @@ public class TcpSendingMessageHandlerTests { assertTrue(replies.remove("Reply1")); assertTrue(replies.remove("Reply2")); done.set(true); + ccf.stop(); } @Test @@ -743,18 +752,19 @@ public class TcpSendingMessageHandlerTests { Executors.newSingleThreadExecutor().execute(new Runnable() { public void run() { try { - ServerSocket server = ServerSocketFactory.getDefault().createServerSocket(port); + ServerSocket server = ServerSocketFactory.getDefault().createServerSocket(port, 100); latch.countDown(); - int i = 0; - while (true) { + for (int i = 0; i < 100; i++) { Socket socket = server.accept(); serverSockets.add(socket); semaphore.release(); byte[] b = new byte[9]; readFully(socket.getInputStream(), b); - b = ("Reply" + (i++) + "\r\n").getBytes(); + b = ("Reply" + i + "\r\n").getBytes(); socket.getOutputStream().write(b); + socket.close(); } + server.close(); } catch (Exception e) { if (!done.get()) { e.printStackTrace(); @@ -797,6 +807,7 @@ public class TcpSendingMessageHandlerTests { assertTrue("Reply" + i + " missing", replies.remove("Reply" + i)); } done.set(true); + ccf.stop(); } @Test