INT-1151 Ensure server side of socket is closed whenever client closes (avoid CLOSE_WAIT status on sockets)

This commit is contained in:
Gary Russell
2010-05-31 17:51:36 +00:00
parent c578fdbbc7
commit 5775fab7b1
8 changed files with 500 additions and 199 deletions

View File

@@ -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;

View File

@@ -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);
}
}

View File

@@ -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)

View File

@@ -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");

View File

@@ -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;

View File

@@ -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;
}

View File

@@ -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());
}
}
}

View File

@@ -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<SelectionKey> keys = selector.selectedKeys();
Iterator<SelectionKey> 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<SelectionKey> keys = selector.selectedKeys();
Iterator<SelectionKey> 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<SelectionKey> keys = selector.selectedKeys();
Iterator<SelectionKey> 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<SelectionKey> keys = selector.selectedKeys();
Iterator<SelectionKey> 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<SelectionKey> keys = selector.selectedKeys();
Iterator<SelectionKey> 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<SelectionKey> keys = selector.selectedKeys();
Iterator<SelectionKey> 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<SelectionKey> keys = selector.selectedKeys();
Iterator<SelectionKey> 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<SelectionKey> keys = selector.selectedKeys();
Iterator<SelectionKey> 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<SelectionKey> keys = selector.selectedKeys();
Iterator<SelectionKey> 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<SelectionKey> keys = selector.selectedKeys();
Iterator<SelectionKey> 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<SelectionKey> keys = selector.selectedKeys();
Iterator<SelectionKey> 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<SelectionKey> keys = selector.selectedKeys();
Iterator<SelectionKey> 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<SelectionKey> keys = selector.selectedKeys();
Iterator<SelectionKey> 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<SelectionKey> keys = selector.selectedKeys();
Iterator<SelectionKey> 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<SelectionKey> keys = selector.selectedKeys();
Iterator<SelectionKey> 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;
}
}