From 5775fab7b1d733238be9bb9a8b307e6251ae00c8 Mon Sep 17 00:00:00 2001 From: Gary Russell Date: Mon, 31 May 2010 17:51:36 +0000 Subject: [PATCH] INT-1151 Ensure server side of socket is closed whenever client closes (avoid CLOSE_WAIT status on sockets) --- .../ip/tcp/AbstractSocketReader.java | 22 +- .../integration/ip/tcp/NetSocketReader.java | 9 +- .../integration/ip/tcp/NioSocketReader.java | 108 +++-- .../ip/tcp/SimpleTcpNetOutboundGateway.java | 5 +- .../integration/ip/tcp/SocketReader.java | 2 + .../ip/tcp/CustomNioSocketReader.java | 7 +- .../ip/tcp/NetSocketReaderTests.java | 158 +++++++ .../ip/tcp/NioSocketReaderTests.java | 388 +++++++++++------- 8 files changed, 500 insertions(+), 199 deletions(-) diff --git a/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/AbstractSocketReader.java b/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/AbstractSocketReader.java index 643cc1c8b4..caf21704ac 100644 --- a/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/AbstractSocketReader.java +++ b/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/AbstractSocketReader.java @@ -91,22 +91,32 @@ public abstract class AbstractSocketReader implements SocketReader, MessageForma protected abstract int assembleDataCustomFormat() throws IOException; public int assembleData() throws IOException { + int result; try { switch (this.messageFormat) { case FORMAT_LENGTH_HEADER: - return assembleDataLengthFormat(); + result = assembleDataLengthFormat(); + break; case FORMAT_STX_ETX: - return assembleDataStxEtxFormat(); + result = assembleDataStxEtxFormat(); + break; case FORMAT_CRLF: - return assembleDataCrLfFormat(); - case FORMAT_CUSTOM: - return assembleDataCustomFormat(); + result = assembleDataCrLfFormat(); + break; case FORMAT_JAVA_SERIALIZED: - return assembleDataSerializedFormat(); + result = assembleDataSerializedFormat(); + break; + case FORMAT_CUSTOM: + result = assembleDataCustomFormat(); + break; default: throw new UnsupportedOperationException( "Unsupported message format: " + messageFormat); } + if (result < 0) { + doClose(); + } + return result; } catch (IOException e) { doClose(); throw e; diff --git a/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/NetSocketReader.java b/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/NetSocketReader.java index 97614ff96d..e159be0efc 100644 --- a/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/NetSocketReader.java +++ b/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/NetSocketReader.java @@ -60,8 +60,9 @@ public class NetSocketReader extends AbstractSocketReader { protected int assembleDataLengthFormat() throws IOException { byte[] lengthPart = new byte[4]; int status = read(lengthPart, true); - if (status < 0) + if (status < 0) { return status; + } int messageLength = ByteBuffer.wrap(lengthPart).getInt(); if (logger.isDebugEnabled()) { logger.debug("Message length is " + messageLength); @@ -147,7 +148,7 @@ public class NetSocketReader extends AbstractSocketReader { } this.assembledData = this.objectInputStream.readObject(); } catch (EOFException ee) { - return -1; + return SOCKET_CLOSED; } catch (ClassNotFoundException e) { throw new IOException(e); } @@ -204,7 +205,9 @@ public class NetSocketReader extends AbstractSocketReader { protected void doClose() { try { socket.close(); - } catch (IOException e) {} + } catch (IOException e) { + logger.error("Error on close", e); + } } diff --git a/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/NioSocketReader.java b/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/NioSocketReader.java index 34e6d37467..4b2f7aa8a6 100644 --- a/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/NioSocketReader.java +++ b/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/NioSocketReader.java @@ -68,7 +68,14 @@ public class NioSocketReader extends AbstractSocketReader { lengthPart = allocate(4); } if (lengthPart.hasRemaining()) { - readChannel(lengthPart); + int status = readChannel(lengthPart); + if (status < 0) { + if (lengthPart.remaining() == 4) { + // not in the middle of a message, clean close + return status; + } + throw new IOException("Channel closed"); + } return MESSAGE_INCOMPLETE; } if (dataPart == null) { @@ -84,7 +91,10 @@ public class NioSocketReader extends AbstractSocketReader { dataPart = ByteBuffer.allocate(messageLength); } if (dataPart.hasRemaining()) { - readChannel(dataPart); + int status = readChannel(dataPart); + if (status < 0) { + throw new IOException("Channel closed"); + } if (dataPart.hasRemaining()) { return MESSAGE_INCOMPLETE; } @@ -99,16 +109,17 @@ public class NioSocketReader extends AbstractSocketReader { */ @Override protected int assembleDataStxEtxFormat() throws IOException { - if (readChannelNonDeterministic()) { - byte bite = rawBuffer.get(); + int len = readChannelNonDeterministic(); + if (len > 0) { + byte bite = this.rawBuffer.get(); int count = 0; - if (!building) { + if (!this.building) { if (bite != STX) { throw new MessageMappingException("Expected STX, received " + Integer.toHexString(bite)); } - building = true; + this.building = true; count++; - if (!rawBuffer.hasRemaining()) { + if (!this.rawBuffer.hasRemaining()) { if (logger.isDebugEnabled()) { logger.debug("Incomplete message, consumed 1 byte"); } @@ -119,27 +130,27 @@ public class NioSocketReader extends AbstractSocketReader { finishAssembly(); return MESSAGE_COMPLETE; } - buildBuffer.put(bite); + this.buildBuffer.put(bite); count++; - if (buildBuffer.position() >= buildBuffer.limit()) { + if (this.buildBuffer.position() >= this.buildBuffer.limit()) { throw new IOException("ETX not found before max message length: " + maxMessageSize); } } while (true) { - if (!rawBuffer.hasRemaining()) { + if (!this.rawBuffer.hasRemaining()) { if (logger.isDebugEnabled()) { logger.debug("Incomplete message, consumed " + count + " bytes"); } return MESSAGE_INCOMPLETE; } - bite = rawBuffer.get(); + bite = this.rawBuffer.get(); if (bite == ETX) { break; } - buildBuffer.put(bite); + this.buildBuffer.put(bite); count++; - if (buildBuffer.position() >= buildBuffer.limit()) { + if (this.buildBuffer.position() >= this.buildBuffer.limit()) { throw new IOException("ETX not found before max message length: " + maxMessageSize); } @@ -149,12 +160,18 @@ public class NioSocketReader extends AbstractSocketReader { } finishAssembly(); return MESSAGE_COMPLETE; + } else if (len == 0) { + logger.debug("Incomplete message, nothing to read"); + return MESSAGE_INCOMPLETE; } else { - if (logger.isDebugEnabled()) { - logger.debug("Incomplete message, consumed 0 bytes"); + logger.debug("Channel closed"); + if (!this.building) { + // not in the middle of a message, clean close + return SOCKET_CLOSED; } + this.building = false; + throw new IOException("Channel closed"); } - return MESSAGE_INCOMPLETE; } /** @@ -162,9 +179,9 @@ public class NioSocketReader extends AbstractSocketReader { */ private void finishAssembly() { byte[] assembledData = new byte[buildBuffer.position()]; - System.arraycopy(buildBuffer.array(), 0, assembledData, 0, assembledData.length); - building = false; - buildBuffer.clear(); + System.arraycopy(this.buildBuffer.array(), 0, assembledData, 0, assembledData.length); + this.building = false; + this.buildBuffer.clear(); this.assembledData = assembledData; logger.debug("Message assembly complete"); } @@ -174,7 +191,8 @@ public class NioSocketReader extends AbstractSocketReader { */ @Override protected int assembleDataCrLfFormat() throws IOException { - if (readChannelNonDeterministic()) { + int len = readChannelNonDeterministic(); + if (len > 0) { int count = 0; while (true) { if (!rawBuffer.hasRemaining()) { @@ -184,18 +202,19 @@ public class NioSocketReader extends AbstractSocketReader { return MESSAGE_INCOMPLETE; } byte bite = rawBuffer.get(); - if (bite == '\n' && buildBuffer.position() > 0) { - buildBuffer.position(buildBuffer.position() - 1); - if (buildBuffer.get() == '\r') { - buildBuffer.position(buildBuffer.position() - 1); + this.building = true; + if (bite == '\n' && this.buildBuffer.position() > 0) { + this.buildBuffer.position(this.buildBuffer.position() - 1); + if (this.buildBuffer.get() == '\r') { + this.buildBuffer.position(this.buildBuffer.position() - 1); break; } } - buildBuffer.put(bite); + this.buildBuffer.put(bite); count++; - if (buildBuffer.position() >= buildBuffer.limit()) { + if (this.buildBuffer.position() >= this.buildBuffer.limit()) { throw new IOException("CRLF not found before max message length: " - + maxMessageSize); + + this.maxMessageSize); } } if (logger.isDebugEnabled()) { @@ -203,12 +222,18 @@ public class NioSocketReader extends AbstractSocketReader { } finishAssembly(); return MESSAGE_COMPLETE; + } else if (len == 0) { + logger.debug("Incomplete message, nothing to read"); + return MESSAGE_INCOMPLETE; } else { - if (logger.isDebugEnabled()) { - logger.debug("Incomplete message, consumed 0 bytes"); + logger.debug("Channel closed"); + if (!this.building) { + // not in the middle of a message, clean close + return SOCKET_CLOSED; } + this.building = false; + throw new IOException("Channel closed"); } - return MESSAGE_INCOMPLETE; } /** @@ -241,18 +266,19 @@ public class NioSocketReader extends AbstractSocketReader { * @param buffer * @throws IOException */ - protected void readChannel(ByteBuffer buffer) throws IOException { + protected int readChannel(ByteBuffer buffer) throws IOException { try { int len = channel.read(buffer); if (len < 0) { logger.debug("Socket closed"); - throw new IOException("Socket closed"); + return len; } if (logger.isDebugEnabled()) { logger.debug("Read " + len + " bytes, buffer is now at " + buffer.position() + " of " + buffer.capacity()); } + return len; } catch (IOException e) { throw e; } @@ -260,10 +286,10 @@ public class NioSocketReader extends AbstractSocketReader { /** * Reads data into the rawBuffer for non-deterministic algorithms. - * @return true if data is available. + * @return bytes remaining in raw buffer or < 0 if channel closed * @throws IOException */ - protected boolean readChannelNonDeterministic() throws IOException { + protected int readChannelNonDeterministic() throws IOException { if (rawBuffer == null) { rawBuffer = allocate(maxMessageSize); buildBuffer = ByteBuffer.allocate(maxMessageSize); @@ -271,22 +297,18 @@ public class NioSocketReader extends AbstractSocketReader { if (logger.isDebugEnabled()) { logger.debug("Raw buffer has " + rawBuffer.remaining() + " remaining"); } - return true; + return rawBuffer.remaining(); } rawBuffer.clear(); int len = channel.read(rawBuffer); - if (len == 0) { - return false; - } if (len < 0) { - logger.debug("Socket closed"); - throw new IOException("Socket closed"); + return len; } rawBuffer.flip(); if (logger.isDebugEnabled()) { logger.debug("Read " + rawBuffer.limit() + " into raw buffer"); } - return true; + return rawBuffer.remaining(); } /** @@ -312,7 +334,9 @@ public class NioSocketReader extends AbstractSocketReader { protected void doClose() { try { channel.close(); - } catch (IOException e) {} + } catch (IOException e) { + logger.error("Error on close", e); + } } /* (non-Javadoc) diff --git a/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/SimpleTcpNetOutboundGateway.java b/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/SimpleTcpNetOutboundGateway.java index 7dc5acae17..1c4621a8c3 100644 --- a/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/SimpleTcpNetOutboundGateway.java +++ b/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/SimpleTcpNetOutboundGateway.java @@ -15,6 +15,7 @@ */ package org.springframework.integration.ip.tcp; +import java.io.IOException; import java.net.Socket; import org.springframework.integration.core.Message; @@ -77,7 +78,9 @@ public class SimpleTcpNetOutboundGateway extends this.soReceiveBufferSize); } try { - this.reader.assembleData(); // Net... always returns true + if (this.reader.assembleData() < 0) { + throw new IOException("Socket closed"); + } Object object = this.reader.getAssembledData(); if (close) { logger.debug("Closing socket because close=true"); diff --git a/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/SocketReader.java b/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/SocketReader.java index 827e603cf9..8144e7d6e7 100644 --- a/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/SocketReader.java +++ b/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/SocketReader.java @@ -30,6 +30,8 @@ import java.net.Socket; */ public interface SocketReader { + public static int SOCKET_CLOSED = -1; + public static int MESSAGE_INCOMPLETE = 0; public static int MESSAGE_COMPLETE = 1; diff --git a/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/CustomNioSocketReader.java b/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/CustomNioSocketReader.java index 88081bb352..272fbbcec9 100644 --- a/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/CustomNioSocketReader.java +++ b/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/CustomNioSocketReader.java @@ -48,7 +48,12 @@ public class CustomNioSocketReader extends NioSocketReader { if (buffer == null) { buffer = allocate(24); } - readChannel(buffer); + int status = readChannel(buffer); + if (status < 0 ) { + if (buffer.remaining() == 24) + return status; + throw new IOException("Channel closed"); + } if (buffer.hasRemaining()) { return MESSAGE_INCOMPLETE; } diff --git a/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/NetSocketReaderTests.java b/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/NetSocketReaderTests.java index b082c11fc3..2fc5abdd6e 100644 --- a/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/NetSocketReaderTests.java +++ b/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/NetSocketReaderTests.java @@ -16,13 +16,17 @@ package org.springframework.integration.ip.tcp; import static org.junit.Assert.assertEquals; +import static org.junit.Assert.assertTrue; import static org.junit.Assert.fail; import java.io.IOException; import java.net.ServerSocket; import java.net.Socket; +import java.util.concurrent.Executors; +import java.util.concurrent.Semaphore; import javax.net.ServerSocketFactory; +import javax.net.SocketFactory; import org.junit.Test; import org.springframework.integration.ip.util.SocketUtils; @@ -271,5 +275,159 @@ public class NetSocketReaderTests { } server.close(); } + + /** + * Tests socket closure when no data received. + * + * @throws Exception + */ + @Test + public void testCloseCleanupNoData() throws Exception { + final int port = SocketUtils.findAvailableServerSocket(); + final Semaphore semaphore = new Semaphore(0); + Executors.newSingleThreadExecutor().execute(new Runnable() { + public void run() { + try { + while (true) { + Socket socket = SocketFactory.getDefault().createSocket("localhost", port); + semaphore.acquire(); + socket.close(); + } + } catch (Exception e) { + e.printStackTrace(); + } + } + }); + try { + ServerSocket server = ServerSocketFactory.getDefault().createServerSocket(port); + server.setSoTimeout(10000); + Socket socket = server.accept(); + NetSocketReader reader = new NetSocketReader(socket); + semaphore.release(); + assertTrue(reader.assembleData() < 0); + assertTrue(reader.getSocket().isClosed()); + + socket = server.accept(); + reader = new NetSocketReader(socket); + reader.setMessageFormat(MessageFormats.FORMAT_CRLF); + semaphore.release(); + assertTrue(reader.assembleData() < 0); + assertTrue(reader.getSocket().isClosed()); + + socket = server.accept(); + reader = new NetSocketReader(socket); + reader.setMessageFormat(MessageFormats.FORMAT_STX_ETX); + semaphore.release(); + assertTrue(reader.assembleData() < 0); + assertTrue(reader.getSocket().isClosed()); + + socket = server.accept(); + reader = new NetSocketReader(socket); + reader.setMessageFormat(MessageFormats.FORMAT_JAVA_SERIALIZED); + semaphore.release(); + assertTrue(reader.assembleData() < 0); + assertTrue(reader.getSocket().isClosed()); + + socket = server.accept(); + reader = new CustomNetSocketReader(socket); + reader.setMessageFormat(MessageFormats.FORMAT_CUSTOM); + semaphore.release(); + assertTrue(reader.assembleData() < 0); + assertTrue(reader.getSocket().isClosed()); + + } catch (IOException e) { + e.printStackTrace(); + fail(e.getMessage()); + } + } + + /** + * Tests socket closure when mid-message + * + * @throws Exception + */ + @Test + public void testCloseCleanup() throws Exception { + final int port = SocketUtils.findAvailableServerSocket(); + final Semaphore semaphore = new Semaphore(0); + Executors.newSingleThreadExecutor().execute(new Runnable() { + public void run() { + try { + Socket socket = SocketFactory.getDefault().createSocket("localhost", port); + byte[] header = {0, 0, 0, 10}; + socket.getOutputStream().write(header); + socket.getOutputStream().write("xx".getBytes()); + semaphore.acquire(); + socket.close(); + + socket = SocketFactory.getDefault().createSocket("localhost", port); + socket.getOutputStream().write("xx".getBytes()); + semaphore.acquire(); + socket.close(); + + socket = SocketFactory.getDefault().createSocket("localhost", port); + socket.getOutputStream().write(MessageFormats.STX); + socket.getOutputStream().write("xx".getBytes()); + semaphore.acquire(); + socket.close(); + + socket = SocketFactory.getDefault().createSocket("localhost", port); + socket.getOutputStream().write(MessageFormats.STX); + socket.getOutputStream().write("xx".getBytes()); + semaphore.acquire(); + socket.close(); + } catch (Exception e) { + e.printStackTrace(); + } + } + }); + try { + ServerSocket server = ServerSocketFactory.getDefault().createServerSocket(port); + server.setSoTimeout(10000); + Socket socket = server.accept(); + NetSocketReader reader = new NetSocketReader(socket); + semaphore.release(); + try { + reader.assembleData(); + fail("Exception expected"); + } catch (IOException e) { } + assertTrue(reader.getSocket().isClosed()); + + socket = server.accept(); + reader = new NetSocketReader(socket); + reader.setMessageFormat(MessageFormats.FORMAT_CRLF); + semaphore.release(); + try { + reader.assembleData(); + fail("Exception expected"); + } catch (IOException e) { } + assertTrue(reader.getSocket().isClosed()); + + socket = server.accept(); + reader = new NetSocketReader(socket); + reader.setMessageFormat(MessageFormats.FORMAT_STX_ETX); + semaphore.release(); + try { + reader.assembleData(); + fail("Exception expected"); + } catch (IOException e) { } + assertTrue(reader.getSocket().isClosed()); + + socket = server.accept(); + reader = new CustomNetSocketReader(socket); + reader.setMessageFormat(MessageFormats.FORMAT_CUSTOM); + semaphore.release(); + try { + reader.assembleData(); + fail("Exception expected"); + } catch (IOException e) { } + assertTrue(reader.getSocket().isClosed()); + + } catch (IOException e) { + e.printStackTrace(); + fail(e.getMessage()); + } + } + } diff --git a/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/NioSocketReaderTests.java b/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/NioSocketReaderTests.java index d9304f55df..8a1f050858 100644 --- a/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/NioSocketReaderTests.java +++ b/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/NioSocketReaderTests.java @@ -21,6 +21,8 @@ import static org.junit.Assert.fail; import java.io.IOException; import java.net.InetSocketAddress; +import java.net.Socket; +import java.nio.channels.ClosedChannelException; import java.nio.channels.SelectionKey; import java.nio.channels.Selector; import java.nio.channels.ServerSocketChannel; @@ -28,6 +30,10 @@ import java.nio.channels.SocketChannel; import java.util.Iterator; import java.util.Set; import java.util.concurrent.CountDownLatch; +import java.util.concurrent.Executors; +import java.util.concurrent.Semaphore; + +import javax.net.SocketFactory; import org.junit.Test; import org.springframework.integration.ip.util.SocketUtils; @@ -52,31 +58,15 @@ public class NioSocketReaderTests { server.register(selector, SelectionKey.OP_ACCEPT); // Fire up the sender. + SocketUtils.testSendLength(port, latch); - if(selector.select(10000) <= 0) { - fail("Socket failed to connect"); - } - Set keys = selector.selectedKeys(); - Iterator iterator = keys.iterator(); - SocketChannel channel = null; - while (iterator.hasNext()) { - SelectionKey key = iterator.next(); - iterator.remove(); - if (key.isAcceptable()) { - channel = server.accept(); - channel.configureBlocking(false); - channel.register(selector, SelectionKey.OP_READ); - } - else { - fail("Unexpected key: " + key); - } - } + SocketChannel channel = accept(server, selector); NioSocketReader reader = new NioSocketReader(channel); int count = 0; while(selector.select(1000) > 0) { - keys = selector.selectedKeys(); - iterator = keys.iterator(); + Set keys = selector.selectedKeys(); + Iterator iterator = keys.iterator(); while (iterator.hasNext()) { SelectionKey key = iterator.next(); iterator.remove(); @@ -109,30 +99,13 @@ public class NioSocketReaderTests { // Fire up the sender. SocketUtils.testSendFragmented(port, false); - - if(selector.select(10000) <= 0) { - fail("Socket failed to connect"); - } - Set keys = selector.selectedKeys(); - Iterator iterator = keys.iterator(); - SocketChannel channel = null; - while (iterator.hasNext()) { - SelectionKey key = iterator.next(); - iterator.remove(); - if (key.isAcceptable()) { - channel = server.accept(); - channel.configureBlocking(false); - channel.register(selector, SelectionKey.OP_READ); - } - else { - fail("Unexpected key: " + key); - } - } + + SocketChannel channel = accept(server, selector); NioSocketReader reader = new NioSocketReader(channel); boolean done = false; while(selector.select(1000) > 0) { - keys = selector.selectedKeys(); - iterator = keys.iterator(); + Set keys = selector.selectedKeys(); + Iterator iterator = keys.iterator(); while (iterator.hasNext()) { SelectionKey key = iterator.next(); iterator.remove(); @@ -168,31 +141,14 @@ public class NioSocketReaderTests { // Fire up the sender. SocketUtils.testSendStxEtx(port, latch); - - if(selector.select(10000) <= 0) { - fail("Socket failed to connect"); - } - Set keys = selector.selectedKeys(); - Iterator iterator = keys.iterator(); - SocketChannel channel = null; - while (iterator.hasNext()) { - SelectionKey key = iterator.next(); - iterator.remove(); - if (key.isAcceptable()) { - channel = server.accept(); - channel.configureBlocking(false); - channel.register(selector, SelectionKey.OP_READ); - } - else { - fail("Unexpected key: " + key); - } - } + + SocketChannel channel = accept(server, selector); NioSocketReader reader = new NioSocketReader(channel); reader.setMessageFormat(MessageFormats.FORMAT_STX_ETX); int count = 0; while(selector.select(1000) > 0) { - keys = selector.selectedKeys(); - iterator = keys.iterator(); + Set keys = selector.selectedKeys(); + Iterator iterator = keys.iterator(); while (iterator.hasNext()) { SelectionKey key = iterator.next(); iterator.remove(); @@ -228,31 +184,14 @@ public class NioSocketReaderTests { // Fire up the sender. SocketUtils.testSendCrLf(port, latch); - - if(selector.select(10000) <= 0) { - fail("Socket failed to connect"); - } - Set keys = selector.selectedKeys(); - Iterator iterator = keys.iterator(); - SocketChannel channel = null; - while (iterator.hasNext()) { - SelectionKey key = iterator.next(); - iterator.remove(); - if (key.isAcceptable()) { - channel = server.accept(); - channel.configureBlocking(false); - channel.register(selector, SelectionKey.OP_READ); - } - else { - fail("Unexpected key: " + key); - } - } + + SocketChannel channel = accept(server, selector); NioSocketReader reader = new NioSocketReader(channel); reader.setMessageFormat(MessageFormats.FORMAT_CRLF); int count = 0; while(selector.select(1000) > 0) { - keys = selector.selectedKeys(); - iterator = keys.iterator(); + Set keys = selector.selectedKeys(); + Iterator iterator = keys.iterator(); while (iterator.hasNext()) { SelectionKey key = iterator.next(); iterator.remove(); @@ -288,30 +227,13 @@ public class NioSocketReaderTests { // Fire up the sender. SocketUtils.testSendLengthOverflow(port); - - if(selector.select(10000) <= 0) { - fail("Socket failed to connect"); - } - Set keys = selector.selectedKeys(); - Iterator iterator = keys.iterator(); - SocketChannel channel = null; - while (iterator.hasNext()) { - SelectionKey key = iterator.next(); - iterator.remove(); - if (key.isAcceptable()) { - channel = server.accept(); - channel.configureBlocking(false); - channel.register(selector, SelectionKey.OP_READ); - } - else { - fail("Unexpected key: " + key); - } - } + + SocketChannel channel = accept(server, selector); NioSocketReader reader = new NioSocketReader(channel); int count = 0; while(selector.select(1000) > 0) { - keys = selector.selectedKeys(); - iterator = keys.iterator(); + Set keys = selector.selectedKeys(); + Iterator iterator = keys.iterator(); while (iterator.hasNext()) { SelectionKey key = iterator.next(); iterator.remove(); @@ -355,32 +277,15 @@ public class NioSocketReaderTests { // Fire up the sender. SocketUtils.testSendStxEtxOverflow(port); - - if(selector.select(10000) <= 0) { - fail("Socket failed to connect"); - } - Set keys = selector.selectedKeys(); - Iterator iterator = keys.iterator(); - SocketChannel channel = null; - while (iterator.hasNext()) { - SelectionKey key = iterator.next(); - iterator.remove(); - if (key.isAcceptable()) { - channel = server.accept(); - channel.configureBlocking(false); - channel.register(selector, SelectionKey.OP_READ); - } - else { - fail("Unexpected key: " + key); - } - } + + SocketChannel channel = accept(server, selector); NioSocketReader reader = new NioSocketReader(channel); reader.setMessageFormat(MessageFormats.FORMAT_STX_ETX); reader.setMaxMessageSize(1024); int count = 0; while(selector.select(1000) > 0) { - keys = selector.selectedKeys(); - iterator = keys.iterator(); + Set keys = selector.selectedKeys(); + Iterator iterator = keys.iterator(); while (iterator.hasNext()) { SelectionKey key = iterator.next(); iterator.remove(); @@ -424,32 +329,15 @@ public class NioSocketReaderTests { // Fire up the sender. SocketUtils.testSendCrLfOverflow(port); - - if(selector.select(10000) <= 0) { - fail("Socket failed to connect"); - } - Set keys = selector.selectedKeys(); - Iterator iterator = keys.iterator(); - SocketChannel channel = null; - while (iterator.hasNext()) { - SelectionKey key = iterator.next(); - iterator.remove(); - if (key.isAcceptable()) { - channel = server.accept(); - channel.configureBlocking(false); - channel.register(selector, SelectionKey.OP_READ); - } - else { - fail("Unexpected key: " + key); - } - } + + SocketChannel channel = accept(server, selector); NioSocketReader reader = new NioSocketReader(channel); reader.setMessageFormat(MessageFormats.FORMAT_CRLF); reader.setMaxMessageSize(1024); int count = 0; while(selector.select(1000) > 0) { - keys = selector.selectedKeys(); - iterator = keys.iterator(); + Set keys = selector.selectedKeys(); + Iterator iterator = keys.iterator(); while (iterator.hasNext()) { SelectionKey key = iterator.next(); iterator.remove(); @@ -479,4 +367,212 @@ public class NioSocketReaderTests { server.close(); } + /** + * Tests socket closure when no data received. + * + * @throws Exception + */ + @Test + public void testCloseCleanupNoData() throws Exception { + final int port = SocketUtils.findAvailableServerSocket(); + final Semaphore semaphore = new Semaphore(0); + Executors.newSingleThreadExecutor().execute(new Runnable() { + public void run() { + try { + semaphore.acquire(); + while (true) { + Socket socket = SocketFactory.getDefault().createSocket("localhost", port); + semaphore.acquire(); + socket.close(); + } + } catch (Exception e) { + e.printStackTrace(); + } + } + }); + try { + ServerSocketChannel server = ServerSocketChannel.open(); + server.configureBlocking(false); + server.socket().bind(new InetSocketAddress(port)); + final Selector selector = Selector.open(); + server.register(selector, SelectionKey.OP_ACCEPT); + + semaphore.release(); + + SocketChannel channel = accept(server, selector); + + NioSocketReader reader = new NioSocketReader(channel); + semaphore.release(); + assertTrue(assembleData(reader) < 0); + assertTrue(reader.getSocket().isClosed()); + + channel = accept(server, selector); + reader = new NioSocketReader(channel); + reader.setMessageFormat(MessageFormats.FORMAT_CRLF); + semaphore.release(); + assertTrue(assembleData(reader) < 0); + assertTrue(reader.getSocket().isClosed()); + + channel = accept(server, selector); + reader = new NioSocketReader(channel); + reader.setMessageFormat(MessageFormats.FORMAT_STX_ETX); + semaphore.release(); + assertTrue(assembleData(reader) < 0); + assertTrue(reader.getSocket().isClosed()); + + channel = accept(server, selector); + reader = new CustomNioSocketReader(channel); + reader.setMessageFormat(MessageFormats.FORMAT_CUSTOM); + semaphore.release(); + assertTrue(assembleData(reader) < 0); + assertTrue(reader.getSocket().isClosed()); + + } catch (IOException e) { + e.printStackTrace(); + fail(e.getMessage()); + } + } + + /** + * Tests socket closure when mid-message + * + * @throws Exception + */ + @Test + public void testCloseCleanup() throws Exception { + final int port = SocketUtils.findAvailableServerSocket(); + final Semaphore semaphore = new Semaphore(0); + Executors.newSingleThreadExecutor().execute(new Runnable() { + public void run() { + try { + semaphore.acquire(); + + Socket socket = SocketFactory.getDefault().createSocket("localhost", port); + byte[] header = {0, 0, 0, 10}; + socket.getOutputStream().write(header); + socket.getOutputStream().write("xx".getBytes()); + semaphore.acquire(); + socket.close(); + + socket = SocketFactory.getDefault().createSocket("localhost", port); + socket.getOutputStream().write("xx".getBytes()); + semaphore.acquire(); + socket.close(); + + socket = SocketFactory.getDefault().createSocket("localhost", port); + socket.getOutputStream().write(MessageFormats.STX); + socket.getOutputStream().write("xx".getBytes()); + semaphore.acquire(); + socket.close(); + + socket = SocketFactory.getDefault().createSocket("localhost", port); + socket.getOutputStream().write(MessageFormats.STX); + socket.getOutputStream().write("xx".getBytes()); + semaphore.acquire(); + socket.close(); + } catch (Exception e) { + e.printStackTrace(); + } + } + }); + try { + ServerSocketChannel server = ServerSocketChannel.open(); + server.configureBlocking(false); + server.socket().bind(new InetSocketAddress(port)); + final Selector selector = Selector.open(); + server.register(selector, SelectionKey.OP_ACCEPT); + + semaphore.release(); + + SocketChannel channel = accept(server, selector); + + NioSocketReader reader = new NioSocketReader(channel); + semaphore.release(); + try { + assembleData(reader); + fail("Exception expected"); + } catch (IOException e) { } + assertTrue(reader.getSocket().isClosed()); + + channel = accept(server, selector); + reader = new NioSocketReader(channel); + reader.setMessageFormat(MessageFormats.FORMAT_CRLF); + semaphore.release(); + try { + assembleData(reader); + fail("Exception expected"); + } catch (IOException e) { } + assertTrue(reader.getSocket().isClosed()); + + channel = accept(server, selector); + reader = new NioSocketReader(channel); + reader.setMessageFormat(MessageFormats.FORMAT_STX_ETX); + semaphore.release(); + try { + assembleData(reader); + fail("Exception expected"); + } catch (IOException e) { } + assertTrue(reader.getSocket().isClosed()); + + channel = accept(server, selector); + reader = new CustomNioSocketReader(channel); + reader.setMessageFormat(MessageFormats.FORMAT_CUSTOM); + semaphore.release(); + try { + assembleData(reader); + fail("Exception expected"); + } catch (IOException e) { } + assertTrue(reader.getSocket().isClosed()); + + } catch (IOException e) { + e.printStackTrace(); + fail(e.getMessage()); + } + } + + + + /** Poor man's nio reader + * + * @param reader + * @return + * @throws IOException + */ + private int assembleData(NioSocketReader reader) throws Exception { + int m = 0; + while (true) { + int n = reader.assembleData(); + if (n < 0) { + return n; + } + Thread.sleep(10); + if (m++ > 1000) + throw new Exception("No close detected"); + } + } + + private SocketChannel accept(ServerSocketChannel server, + final Selector selector) throws IOException, ClosedChannelException { + SocketChannel channel = null; + + if(selector.select(10000) <= 0) { + fail("Socket failed to connect"); + } + Set keys = selector.selectedKeys(); + Iterator iterator = keys.iterator(); + while (iterator.hasNext()) { + SelectionKey key = iterator.next(); + iterator.remove(); + if (key.isAcceptable()) { + channel = server.accept(); + channel.configureBlocking(false); + channel.register(selector, SelectionKey.OP_READ); + } + else { + fail("Unexpected key: " + key); + } + } + return channel; + } + }