diff --git a/org.springframework.integration.ip/src/main/java/org/springframework/integration/ip/tcp/AbstractSocketReader.java b/org.springframework.integration.ip/src/main/java/org/springframework/integration/ip/tcp/AbstractSocketReader.java index dcd5959b56..1bfe795ba1 100644 --- a/org.springframework.integration.ip/src/main/java/org/springframework/integration/ip/tcp/AbstractSocketReader.java +++ b/org.springframework.integration.ip/src/main/java/org/springframework/integration/ip/tcp/AbstractSocketReader.java @@ -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; + } } diff --git a/org.springframework.integration.ip/src/main/java/org/springframework/integration/ip/tcp/NetSocketReader.java b/org.springframework.integration.ip/src/main/java/org/springframework/integration/ip/tcp/NetSocketReader.java index 241d74a309..63fa123164 100644 --- a/org.springframework.integration.ip/src/main/java/org/springframework/integration/ip/tcp/NetSocketReader.java +++ b/org.springframework.integration.ip/src/main/java/org/springframework/integration/ip/tcp/NetSocketReader.java @@ -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); diff --git a/org.springframework.integration.ip/src/main/java/org/springframework/integration/ip/tcp/NioSocketReader.java b/org.springframework.integration.ip/src/main/java/org/springframework/integration/ip/tcp/NioSocketReader.java index 45a9993419..8e33c251bf 100644 --- a/org.springframework.integration.ip/src/main/java/org/springframework/integration/ip/tcp/NioSocketReader.java +++ b/org.springframework.integration.ip/src/main/java/org/springframework/integration/ip/tcp/NioSocketReader.java @@ -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"); diff --git a/org.springframework.integration.ip/src/main/java/org/springframework/integration/ip/tcp/TcpNetReceivingChannelAdapter.java b/org.springframework.integration.ip/src/main/java/org/springframework/integration/ip/tcp/TcpNetReceivingChannelAdapter.java index e5463bc399..7dc733b5b2 100644 --- a/org.springframework.integration.ip/src/main/java/org/springframework/integration/ip/tcp/TcpNetReceivingChannelAdapter.java +++ b/org.springframework.integration.ip/src/main/java/org/springframework/integration/ip/tcp/TcpNetReceivingChannelAdapter.java @@ -107,6 +107,7 @@ public class TcpNetReceivingChannelAdapter extends reader = new NetSocketReader(socket); } reader.setMessageFormat(messageFormat); + reader.setMaxMessageSize(receiveBufferSize); while (true) { try { if (reader.assembleData()) { diff --git a/org.springframework.integration.ip/src/main/java/org/springframework/integration/ip/tcp/TcpNioReceivingChannelAdapter.java b/org.springframework.integration.ip/src/main/java/org/springframework/integration/ip/tcp/TcpNioReceivingChannelAdapter.java index c4a60bbca5..456a784a22 100644 --- a/org.springframework.integration.ip/src/main/java/org/springframework/integration/ip/tcp/TcpNioReceivingChannelAdapter.java +++ b/org.springframework.integration.ip/src/main/java/org/springframework/integration/ip/tcp/TcpNioReceivingChannelAdapter.java @@ -167,6 +167,7 @@ public class TcpNioReceivingChannelAdapter extends } reader.setUsingDirectBuffers(usingDirectBuffers); reader.setMessageFormat(messageFormat); + reader.setMaxMessageSize(receiveBufferSize); return reader; } diff --git a/org.springframework.integration.ip/src/test/java/org/springframework/integration/ip/tcp/NetSocketReaderTests.java b/org.springframework.integration.ip/src/test/java/org/springframework/integration/ip/tcp/NetSocketReaderTests.java index 3e8b3ccd3f..5aabc4f6d9 100644 --- a/org.springframework.integration.ip/src/test/java/org/springframework/integration/ip/tcp/NetSocketReaderTests.java +++ b/org.springframework.integration.ip/src/test/java/org/springframework/integration/ip/tcp/NetSocketReaderTests.java @@ -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(); + } + } diff --git a/org.springframework.integration.ip/src/test/java/org/springframework/integration/ip/tcp/NioSocketReaderTests.java b/org.springframework.integration.ip/src/test/java/org/springframework/integration/ip/tcp/NioSocketReaderTests.java index da2f97d16c..a0e72edbe0 100644 --- a/org.springframework.integration.ip/src/test/java/org/springframework/integration/ip/tcp/NioSocketReaderTests.java +++ b/org.springframework.integration.ip/src/test/java/org/springframework/integration/ip/tcp/NioSocketReaderTests.java @@ -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 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); + } + } + 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 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); + } + } + 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 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); + } + } + 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(); + } + } diff --git a/org.springframework.integration.ip/src/test/java/org/springframework/integration/ip/util/SocketUtils.java b/org.springframework.integration.ip/src/test/java/org/springframework/integration/ip/util/SocketUtils.java index c6e1c80872..c44b26fae0 100644 --- a/org.springframework.integration.ip/src/test/java/org/springframework/integration/ip/util/SocketUtils.java +++ b/org.springframework.integration.ip/src/test/java/org/springframework/integration/ip/util/SocketUtils.java @@ -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 {