TCP: Add length check against receive buffer size to prevent OOM error when receiving a message with bogus length. Also test for overflow on other formats.

This commit is contained in:
Gary Russell
2010-03-16 16:20:36 +00:00
parent efe0f1e54d
commit 44608ace66
8 changed files with 428 additions and 18 deletions

View File

@@ -44,6 +44,8 @@ public abstract class AbstractSocketReader implements SocketReader, MessageForma
* returns true; will be set to null when getAssembledData() is called.
*/
protected byte[] assembledData;
protected int maxMessageSize = 1024 * 60;
/**
* Assembles data in format {@link #FORMAT_LENGTH_HEADER}.
@@ -109,4 +111,7 @@ public abstract class AbstractSocketReader implements SocketReader, MessageForma
this.messageFormat = messageFormat;
}
public void setMaxMessageSize(int maxMessageSize) {
this.maxMessageSize = maxMessageSize;
}
}

View File

@@ -39,8 +39,6 @@ public class NetSocketReader extends AbstractSocketReader {
protected Socket socket;
protected int receiveBufferSize = 1024 * 60;
/**
* Constructs a NetsocketReader which reads from the Socket.
* @param socket The socket.
@@ -59,7 +57,11 @@ public class NetSocketReader extends AbstractSocketReader {
int messageLength = ByteBuffer.wrap(lengthPart).getInt();
if (logger.isDebugEnabled()) {
logger.debug("Message length is " + messageLength);
}
}
if (messageLength > maxMessageSize) {
throw new IOException("Message length " + messageLength +
" exceeds max message length: " + maxMessageSize);
}
byte[] messagePart = new byte[messageLength];
read(messagePart);
assembledData = messagePart;
@@ -74,11 +76,15 @@ public class NetSocketReader extends AbstractSocketReader {
InputStream inputStream = socket.getInputStream();
if (inputStream.read() != STX)
throw new MessageMappingException("Expected STX to begin message");
byte[] buffer = new byte[receiveBufferSize];
byte[] buffer = new byte[maxMessageSize];
int n = 0;
int bite;
while ((bite = inputStream.read()) != ETX) {
buffer[n++] = (byte) bite;
if (n >= maxMessageSize) {
throw new IOException("ETX not found before max message length: "
+ maxMessageSize);
}
}
assembledData = new byte[n];
System.arraycopy(buffer, 0, assembledData, 0, n);
@@ -91,7 +97,7 @@ public class NetSocketReader extends AbstractSocketReader {
@Override
protected boolean assembleDataCrLfFormat() throws IOException {
InputStream inputStream = socket.getInputStream();
byte[] buffer = new byte[receiveBufferSize];
byte[] buffer = new byte[maxMessageSize];
int n = 0;
int bite;
while (true) {
@@ -99,6 +105,10 @@ public class NetSocketReader extends AbstractSocketReader {
if (n > 0 && bite == '\n' && buffer[n-1] == '\r')
break;
buffer[n++] = (byte) bite;
if (n >= maxMessageSize) {
throw new IOException("CRLF not found before max message length: "
+ maxMessageSize);
}
};
assembledData = new byte[n-1];
System.arraycopy(buffer, 0, assembledData, 0, n-1);

View File

@@ -46,8 +46,6 @@ public class NioSocketReader extends AbstractSocketReader {
protected ByteBuffer buildBuffer;
protected int receiveBufferSize = 1024 * 60;
protected boolean building;
/**
@@ -85,7 +83,11 @@ public class NioSocketReader extends AbstractSocketReader {
if (logger.isDebugEnabled()) {
logger.debug("Message length is " + messageLength);
}
dataPart = allocate(messageLength);
if (messageLength > maxMessageSize) {
throw new IOException("Message length " + messageLength +
" exceeds max message length " + maxMessageSize);
}
dataPart = ByteBuffer.allocate(messageLength);
}
if (dataPart.hasRemaining()) {
readChannel(dataPart);
@@ -93,14 +95,7 @@ public class NioSocketReader extends AbstractSocketReader {
return false;
}
}
if (usingDirectBuffers) {
byte[] assembledData = new byte[dataPart.capacity()];
dataPart.flip();
dataPart.get(assembledData);
this.assembledData = assembledData;
} else {
assembledData = dataPart.array();
}
assembledData = dataPart.array();
lengthPart = dataPart = null;
return true;
}
@@ -132,6 +127,10 @@ public class NioSocketReader extends AbstractSocketReader {
}
buildBuffer.put(bite);
count++;
if (buildBuffer.position() >= buildBuffer.limit()) {
throw new IOException("ETX not found before max message length: "
+ maxMessageSize);
}
}
while (true) {
if (!rawBuffer.hasRemaining()) {
@@ -146,6 +145,10 @@ public class NioSocketReader extends AbstractSocketReader {
}
buildBuffer.put(bite);
count++;
if (buildBuffer.position() >= buildBuffer.limit()) {
throw new IOException("ETX not found before max message length: "
+ maxMessageSize);
}
}
if (logger.isDebugEnabled()) {
logger.debug("Consumed " + count + " bytes");
@@ -195,6 +198,10 @@ public class NioSocketReader extends AbstractSocketReader {
}
buildBuffer.put(bite);
count++;
if (buildBuffer.position() >= buildBuffer.limit()) {
throw new IOException("CRLF not found before max message length: "
+ maxMessageSize);
}
}
if (logger.isDebugEnabled()) {
logger.debug("Consumed " + count + " bytes");
@@ -251,8 +258,8 @@ public class NioSocketReader extends AbstractSocketReader {
*/
protected boolean readChannelNonDeterministic() throws IOException {
if (rawBuffer == null) {
rawBuffer = allocate(receiveBufferSize);
buildBuffer = ByteBuffer.allocate(receiveBufferSize);
rawBuffer = allocate(maxMessageSize);
buildBuffer = ByteBuffer.allocate(maxMessageSize);
} else if (rawBuffer.hasRemaining()) {
if (logger.isDebugEnabled()) {
logger.debug("Raw buffer has " + rawBuffer.remaining() + " remaining");

View File

@@ -107,6 +107,7 @@ public class TcpNetReceivingChannelAdapter extends
reader = new NetSocketReader(socket);
}
reader.setMessageFormat(messageFormat);
reader.setMaxMessageSize(receiveBufferSize);
while (true) {
try {
if (reader.assembleData()) {

View File

@@ -167,6 +167,7 @@ public class TcpNioReceivingChannelAdapter extends
}
reader.setUsingDirectBuffers(usingDirectBuffers);
reader.setMessageFormat(messageFormat);
reader.setMaxMessageSize(receiveBufferSize);
return reader;
}

View File

@@ -16,8 +16,10 @@
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;
@@ -121,4 +123,115 @@ public class NetSocketReaderTests {
server.close();
}
@Test
public void testReadLengthOverflow() throws Exception {
int port = SocketUtils.findAvailableServerSocket();
ServerSocket server = ServerSocketFactory.getDefault().createServerSocket(port);
SocketUtils.testSendLengthOverflow(port);
Socket socket = server.accept();
socket.setSoTimeout(5000);
NetSocketReader reader = new NetSocketReader(socket);
try {
if (reader.assembleData()) {
fail("Expected message length exceeded exception");
}
} catch (IOException e) {
if (!e.getMessage().startsWith("Message length")) {
e.printStackTrace();
fail("Unexpected IO Error:" + e.getMessage());
}
}
server.close();
}
@Test
public void testReadStxEtxTimeout() throws Exception {
int port = SocketUtils.findAvailableServerSocket();
ServerSocket server = ServerSocketFactory.getDefault().createServerSocket(port);
SocketUtils.testSendStxEtxOverflow(port);
Socket socket = server.accept();
socket.setSoTimeout(500);
NetSocketReader reader = new NetSocketReader(socket);
reader.setMessageFormat(MessageFormats.FORMAT_STX_ETX);
try {
if (reader.assembleData()) {
fail("Expected message length exceeded exception");
}
} catch (IOException e) {
if (!e.getMessage().startsWith("Read timed out")) {
e.printStackTrace();
fail("Unexpected IO Error:" + e.getMessage());
}
}
server.close();
}
@Test
public void testReadStxEtxOverflow() throws Exception {
int port = SocketUtils.findAvailableServerSocket();
ServerSocket server = ServerSocketFactory.getDefault().createServerSocket(port);
SocketUtils.testSendStxEtxOverflow(port);
Socket socket = server.accept();
socket.setSoTimeout(5000);
NetSocketReader reader = new NetSocketReader(socket);
reader.setMessageFormat(MessageFormats.FORMAT_STX_ETX);
reader.setMaxMessageSize(1024);
try {
if (reader.assembleData()) {
fail("Expected message length exceeded exception");
}
} catch (IOException e) {
if (!e.getMessage().startsWith("ETX not found")) {
e.printStackTrace();
fail("Unexpected IO Error:" + e.getMessage());
}
}
server.close();
}
@Test
public void testReadCrLfTimeout() throws Exception {
int port = SocketUtils.findAvailableServerSocket();
ServerSocket server = ServerSocketFactory.getDefault().createServerSocket(port);
SocketUtils.testSendCrLfOverflow(port);
Socket socket = server.accept();
socket.setSoTimeout(500);
NetSocketReader reader = new NetSocketReader(socket);
reader.setMessageFormat(MessageFormats.FORMAT_CRLF);
try {
if (reader.assembleData()) {
fail("Expected message length exceeded exception");
}
} catch (IOException e) {
if (!e.getMessage().startsWith("Read timed out")) {
e.printStackTrace();
fail("Unexpected IO Error:" + e.getMessage());
}
}
server.close();
}
@Test
public void testReadCrLfOverflow() throws Exception {
int port = SocketUtils.findAvailableServerSocket();
ServerSocket server = ServerSocketFactory.getDefault().createServerSocket(port);
SocketUtils.testSendCrLfOverflow(port);
Socket socket = server.accept();
socket.setSoTimeout(5000);
NetSocketReader reader = new NetSocketReader(socket);
reader.setMessageFormat(MessageFormats.FORMAT_CRLF);
reader.setMaxMessageSize(1024);
try {
if (reader.assembleData()) {
fail("Expected message length exceeded exception");
}
} catch (IOException e) {
if (!e.getMessage().startsWith("CRLF not found")) {
e.printStackTrace();
fail("Unexpected IO Error:" + e.getMessage());
}
}
server.close();
}
}

View File

@@ -19,6 +19,7 @@ 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.InetSocketAddress;
import java.nio.channels.SelectionKey;
import java.nio.channels.Selector;
@@ -273,4 +274,209 @@ public class NioSocketReaderTests {
server.close();
}
/**
* Test method for {@link org.springframework.integration.ip.tcp.NioSocketReader}.
*/
@Test
public void testReadLengthOverflow() throws Exception {
ServerSocketChannel server = ServerSocketChannel.open();
server.configureBlocking(false);
int port = SocketUtils.findAvailableServerSocket();
server.socket().bind(new InetSocketAddress(port));
final Selector selector = Selector.open();
server.register(selector, SelectionKey.OP_ACCEPT);
// 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);
}
}
NioSocketReader reader = new NioSocketReader(channel);
int count = 0;
while(selector.select(1000) > 0) {
keys = selector.selectedKeys();
iterator = keys.iterator();
while (iterator.hasNext()) {
SelectionKey key = iterator.next();
iterator.remove();
if (key.isReadable()) {
assertEquals(channel, key.channel());
try {
if (reader.assembleData()) {
fail("Expected message length exceeded exception");
}
} catch (IOException e) {
if (!e.getMessage().startsWith("Message length")) {
e.printStackTrace();
fail("Unexpected IO Error:" + e.getMessage());
}
count++;
break;
}
}
else {
fail("Unexpected key: " + key);
}
}
if (count > 0) {
break;
}
}
server.close();
}
/**
* Test method for {@link org.springframework.integration.ip.tcp.NioSocketReader}.
*/
@Test
public void testReadStxEtxOverflow() throws Exception {
ServerSocketChannel server = ServerSocketChannel.open();
server.configureBlocking(false);
int port = SocketUtils.findAvailableServerSocket();
server.socket().bind(new InetSocketAddress(port));
final Selector selector = Selector.open();
server.register(selector, SelectionKey.OP_ACCEPT);
// 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);
}
}
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();
while (iterator.hasNext()) {
SelectionKey key = iterator.next();
iterator.remove();
if (key.isReadable()) {
assertEquals(channel, key.channel());
try {
if (reader.assembleData()) {
fail("Expected message length exceeded exception");
}
} catch (IOException e) {
if (!e.getMessage().startsWith("ETX not found")) {
e.printStackTrace();
fail("Unexpected IO Error:" + e.getMessage());
}
count++;
break;
}
}
else {
fail("Unexpected key: " + key);
}
}
if (count > 0) {
break;
}
}
server.close();
}
/**
* Test method for {@link org.springframework.integration.ip.tcp.NioSocketReader}.
*/
@Test
public void testReadCrLfOverflow() throws Exception {
ServerSocketChannel server = ServerSocketChannel.open();
server.configureBlocking(false);
int port = SocketUtils.findAvailableServerSocket();
server.socket().bind(new InetSocketAddress(port));
final Selector selector = Selector.open();
server.register(selector, SelectionKey.OP_ACCEPT);
// 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);
}
}
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();
while (iterator.hasNext()) {
SelectionKey key = iterator.next();
iterator.remove();
if (key.isReadable()) {
assertEquals(channel, key.channel());
try {
if (reader.assembleData()) {
fail("Expected message length exceeded exception");
}
} catch (IOException e) {
if (!e.getMessage().startsWith("CRLF not found")) {
e.printStackTrace();
fail("Unexpected IO Error:" + e.getMessage());
}
count++;
break;
}
}
else {
fail("Unexpected key: " + key);
}
}
if (count > 0) {
break;
}
}
server.close();
}
}

View File

@@ -73,6 +73,28 @@ public class SocketUtils {
thread.start();
}
/**
* Sends a message with a bad length part, causing an overflow on the receiver.
*/
public static void testSendLengthOverflow(final int port) {
Thread thread = new Thread(new Runnable() {
public void run() {
try {
Socket socket = new Socket(InetAddress.getByName("localhost"), port);
byte[] len = new byte[4];
ByteBuffer.wrap(len).putInt(Integer.MAX_VALUE);
socket.getOutputStream().write(len);
socket.getOutputStream().write(TEST_STRING.getBytes());
Thread.sleep(1000000000L); // wait forever, but we're a daemon
} catch (Exception e) {
e.printStackTrace();
}
}
});
thread.setDaemon(true);
thread.start();
}
/**
* Test for reassembly of completely fragmented message; sends
* 6 bytes 500ms apart.
@@ -145,6 +167,29 @@ public class SocketUtils {
thread.start();
}
/**
* Sends a large STX/ETX message with no ETX
*/
public static void testSendStxEtxOverflow(final int port) {
Thread thread = new Thread(new Runnable() {
public void run() {
try {
Socket socket = new Socket(InetAddress.getByName("localhost"), port);
OutputStream outputStream = socket.getOutputStream();
writeByte(outputStream, 0x02, true);
for (int i = 0; i < 1500; i++) {
writeByte(outputStream, 'x', true);
}
Thread.sleep(1000000000L); // wait forever, but we're a daemon
} catch (Exception e) {
e.printStackTrace();
}
}
});
thread.setDaemon(true);
thread.start();
}
/**
* Sends a message +CRLF in two chunks. Two such messages are sent.
* @param latch If not null, await until counted down before sending second chunk.
@@ -178,6 +223,28 @@ public class SocketUtils {
thread.start();
}
/**
* Sends a large CRLF message with no CRLF.
*/
public static void testSendCrLfOverflow(final int port) {
Thread thread = new Thread(new Runnable() {
public void run() {
try {
Socket socket = new Socket(InetAddress.getByName("localhost"), port);
OutputStream outputStream = socket.getOutputStream();
for (int i = 0; i < 1500; i++) {
writeByte(outputStream, 'x', true);
}
Thread.sleep(1000000000L); // wait forever, but we're a daemon
} catch (Exception e) {
e.printStackTrace();
}
}
});
thread.setDaemon(true);
thread.start();
}
public static int findAvailableServerSocket(int seed) {
for (int i = seed; i < seed+200; i++) {
try {