From 86fa138554aa9e626df0077b529cc085c4667e35 Mon Sep 17 00:00:00 2001 From: Gary Russell Date: Fri, 27 Jun 2014 18:22:07 +0300 Subject: [PATCH] INT-3453 NIO Always Release Assembler Thread Fixed hung assembler described in INT-3453. Used CountDownLatch instead of sleep() INT-3453 Add Test Case Currently failing for me - assembler thread stuck on CDL. Remove redundant `if (TcpNioConnection.this.writingLatch != null)` from `ChannelInputStream#write` --- .../ip/tcp/connection/TcpNioConnection.java | 36 ++++++- .../tcp/connection/TcpNioConnectionTests.java | 99 +++++++++++++++++++ 2 files changed, 132 insertions(+), 3 deletions(-) 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 a42a03ea29..8d236d1de1 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 @@ -26,6 +26,7 @@ import java.nio.channels.SelectionKey; import java.nio.channels.Selector; import java.nio.channels.SocketChannel; import java.util.concurrent.BlockingQueue; +import java.util.concurrent.CountDownLatch; import java.util.concurrent.Executor; import java.util.concurrent.ExecutorService; import java.util.concurrent.Executors; @@ -75,6 +76,8 @@ public class TcpNioConnection extends TcpConnectionSupport { private volatile boolean writingToPipe; + private volatile CountDownLatch writingLatch; + private volatile long pipeTimeout = DEFAULT_PIPE_TIMEOUT; /** @@ -285,7 +288,11 @@ public class TcpNioConnection extends TcpConnectionSupport { } private boolean dataAvailable() throws IOException { - return this.channelInputStream.available() > 0 || writingToPipe; + if (logger.isTraceEnabled()) { + logger.trace(getConnectionId() + " checking data avail: " + this.channelInputStream.available() + + " pending: " + (this.writingToPipe)); + } + return writingToPipe || this.channelInputStream.available() > 0; } /** @@ -295,8 +302,25 @@ public class TcpNioConnection extends TcpConnectionSupport { * @throws IOException */ private synchronized Message convert() throws Exception { - if (!dataAvailable()) { - return null; + if (logger.isTraceEnabled()) { + logger.trace(getConnectionId() + " checking data avail (convert): " + this.channelInputStream.available() + + " pending: " + (this.writingToPipe)); + } + if (this.channelInputStream.available() <= 0) { + try { + if (this.writingLatch.await(60, TimeUnit.SECONDS)) { + if (this.channelInputStream.available() <= 0) { + return null; + } + } + else { // should never happen + throw new IOException("Timed out waiting for IO"); + } + } + catch (InterruptedException e) { + Thread.currentThread().interrupt(); + throw new IOException("Interrupted waiting for IO"); + } } Message message = null; try { @@ -356,6 +380,7 @@ public class TcpNioConnection extends TcpConnectionSupport { this.rawBuffer = allocate(maxMessageSize); } + this.writingLatch = new CountDownLatch(1); this.writingToPipe = true; try { if (this.taskExecutor == null) { @@ -394,6 +419,7 @@ public class TcpNioConnection extends TcpConnectionSupport { } finally { this.writingToPipe = false; + this.writingLatch.countDown(); } } @@ -682,6 +708,9 @@ public class TcpNioConnection extends TcpConnectionSupport { byte[] buffer = new byte[bytesToWrite]; System.arraycopy(array, 0, buffer, 0, bytesToWrite); this.available.addAndGet(bytesToWrite); + if (TcpNioConnection.this.writingLatch != null) { + TcpNioConnection.this.writingLatch.countDown(); + } try { if (!this.buffers.offer(buffer, pipeTimeout, TimeUnit.MILLISECONDS)) { throw new IOException("Timed out waiting for buffer space"); @@ -691,6 +720,7 @@ public class TcpNioConnection extends TcpConnectionSupport { Thread.currentThread().interrupt(); throw new IOException("Interrupted while waiting for buffer space", e); } + TcpNioConnection.this.writingLatch = new CountDownLatch(1); } } diff --git a/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/connection/TcpNioConnectionTests.java b/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/connection/TcpNioConnectionTests.java index e704eb2dc8..e32b27db66 100644 --- a/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/connection/TcpNioConnectionTests.java +++ b/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/connection/TcpNioConnectionTests.java @@ -17,6 +17,7 @@ package org.springframework.integration.ip.tcp.connection; import static org.hamcrest.Matchers.containsString; +import static org.hamcrest.Matchers.not; import static org.junit.Assert.assertEquals; import static org.junit.Assert.assertNotNull; import static org.junit.Assert.assertThat; @@ -24,7 +25,9 @@ import static org.junit.Assert.assertTrue; import static org.junit.Assert.fail; import static org.mockito.Matchers.any; import static org.mockito.Mockito.doAnswer; +import static org.mockito.Mockito.doReturn; import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.spy; import static org.mockito.Mockito.when; import java.io.ByteArrayOutputStream; @@ -41,6 +44,7 @@ import java.nio.channels.SelectionKey; import java.nio.channels.Selector; import java.nio.channels.SocketChannel; import java.util.ArrayList; +import java.util.Arrays; import java.util.HashMap; import java.util.HashSet; import java.util.List; @@ -58,7 +62,9 @@ import java.util.concurrent.atomic.AtomicReference; import javax.net.ServerSocketFactory; import javax.net.SocketFactory; +import org.apache.commons.logging.Log; import org.junit.Test; +import org.mockito.Matchers; import org.mockito.Mockito; import org.mockito.invocation.InvocationOnMock; import org.mockito.stubbing.Answer; @@ -69,6 +75,7 @@ import org.springframework.context.ApplicationEventPublisher; import org.springframework.integration.ip.tcp.connection.TcpNioConnection.ChannelInputStream; import org.springframework.integration.ip.tcp.serializer.ByteArrayCrLfSerializer; import org.springframework.integration.ip.tcp.serializer.MapJsonSerializer; +import org.springframework.integration.ip.util.TestingUtilities; import org.springframework.integration.support.MessageBuilder; import org.springframework.integration.support.converter.MapMessageConverter; import org.springframework.integration.test.util.SocketUtils; @@ -695,6 +702,98 @@ public class TcpNioConnectionTests { return new CompositeExecutor(ioExec, assemblerExec); } + @Test + public void int3453RaceTest() throws Exception { + final int port = SocketUtils.findAvailableServerSocket(); + TcpNioServerConnectionFactory factory = new TcpNioServerConnectionFactory(port); + final CountDownLatch connectionLatch = new CountDownLatch(1); + factory.setApplicationEventPublisher(new ApplicationEventPublisher() { + + @Override + public void publishEvent(ApplicationEvent event) { + connectionLatch.countDown(); + } + }); + final CountDownLatch assemblerLatch = new CountDownLatch(1); + final AtomicReference assembler = new AtomicReference(); + factory.registerListener(new TcpListener() { + + @Override + public boolean onMessage(Message message) { + if (!(message instanceof ErrorMessage)) { + assemblerLatch.countDown(); + assembler.set(Thread.currentThread()); + } + return false; + } + + }); + ThreadPoolTaskExecutor te = new ThreadPoolTaskExecutor(); + te.setCorePoolSize(3); // selector, reader, assembler + te.setMaxPoolSize(3); + te.setQueueCapacity(0); + te.initialize(); + factory.setTaskExecutor(te); + factory.start(); + TestingUtilities.waitListening(factory, 10000L); + Socket socket = SocketFactory.getDefault().createSocket("localhost", port); + assertTrue(connectionLatch.await(10, TimeUnit.SECONDS)); + + TcpNioConnection connection = (TcpNioConnection) TestUtils.getPropertyValue(factory, "connections", List.class).get(0); + Log logger = spy(TestUtils.getPropertyValue(connection, "logger", Log.class)); + DirectFieldAccessor dfa = new DirectFieldAccessor(connection); + dfa.setPropertyValue("logger", logger); + + ChannelInputStream cis = spy(TestUtils.getPropertyValue(connection, "channelInputStream", ChannelInputStream.class)); + dfa.setPropertyValue("channelInputStream", cis); + + final CountDownLatch readerLatch = new CountDownLatch(4); // 3 dataAvailable, 1 continuing + final CountDownLatch readerFinishedLatch = new CountDownLatch(1); + doAnswer(new Answer() { + + @Override + public Void answer(InvocationOnMock invocation) throws Throwable { + invocation.callRealMethod(); + // delay the reader thread resetting writingToPipe + readerLatch.await(10, TimeUnit.SECONDS); + Thread.sleep(100); + readerFinishedLatch.countDown(); + return null; + } + }).when(cis).write(any(byte[].class), Matchers.anyInt()); + + doReturn(true).when(logger).isTraceEnabled(); + doAnswer(new Answer() { + + @Override + public Void answer(InvocationOnMock invocation) throws Throwable { + invocation.callRealMethod(); + readerLatch.countDown(); + return null; + } + }).when(logger).trace(Matchers.contains("checking data avail")); + + doAnswer(new Answer() { + + @Override + public Void answer(InvocationOnMock invocation) throws Throwable { + invocation.callRealMethod(); + readerLatch.countDown(); + return null; + } + }).when(logger).trace(Matchers.contains("Nio assembler continuing")); + + socket.getOutputStream().write("foo\r\n".getBytes()); + + assertTrue(assemblerLatch.await(10, TimeUnit.SECONDS)); + assertTrue(readerFinishedLatch.await(10, TimeUnit.SECONDS)); + + StackTraceElement[] stackTrace = assembler.get().getStackTrace(); + assertThat(Arrays.asList(stackTrace).toString(), not(containsString("ChannelInputStream.getNextBuffer"))); + socket.close(); + factory.stop(); + } + private void readFully(InputStream is, byte[] buff) throws IOException { for (int i = 0; i < buff.length; i++) { buff[i] = (byte) is.read();