From 474a00509bdce5113ad42d4ab29d6bd95ac1009f Mon Sep 17 00:00:00 2001 From: Gary Russell Date: Mon, 11 Mar 2013 22:34:14 -0400 Subject: [PATCH] INT-2956 Fix NIO With Custom Deserializers The TcpNioConnection uses an internal InputStream. The standard deserializers only use the read() method. Custom Deserializers might use read(byte[]); the internal InputStream did not override these methods, possibly causing the deserializer to hang awaiting data, when none was coming. Override these methods, with the appropriate semantics from the super class... - Block until at least one byte arrives. - Return early if no data is available after at least one byte arrives. - Return -1 if closed when no data arrived. - Return the number of bytes actually read. Polishing - PR Comments remove read(byte[]) because it doesn't do anything different to the superclass method. --- .../ip/tcp/connection/TcpNioConnection.java | 27 +++++ .../tcp/connection/TcpNioConnectionTests.java | 105 ++++++++++++++++++ 2 files changed, 132 insertions(+) diff --git a/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/connection/TcpNioConnection.java b/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/connection/TcpNioConnection.java index b368a70742..a4bf9a7f5f 100644 --- a/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/connection/TcpNioConnection.java +++ b/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/connection/TcpNioConnection.java @@ -547,6 +547,33 @@ public class TcpNioConnection extends TcpConnectionSupport { private volatile boolean isClosed; + @Override + public int read(byte[] b, int off, int len) throws IOException { + Assert.notNull(b, "byte[] cannot be null"); + if (off < 0 || len < 0 || len > b.length - off) { + throw new IndexOutOfBoundsException(); + } + else if (len == 0) { + return 0; + } + + int n = 0; + while ((this.available.get() > 0 || n == 0) && + n < len) { + int bite = read(); + if (bite < 0) { + if (n == 0) { + return -1; + } + else { + return n; + } + } + b[off + n++] = (byte) bite; + } + return n; + } + @Override public synchronized int read() throws IOException { if (this.isClosed && available.get() == 0) { diff --git a/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/connection/TcpNioConnectionTests.java b/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/connection/TcpNioConnectionTests.java index 590351f157..9da0c95fa7 100644 --- a/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/connection/TcpNioConnectionTests.java +++ b/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/connection/TcpNioConnectionTests.java @@ -53,7 +53,9 @@ import org.junit.Test; import org.mockito.Mockito; import org.mockito.invocation.InvocationOnMock; import org.mockito.stubbing.Answer; +import org.springframework.beans.DirectFieldAccessor; import org.springframework.integration.Message; +import org.springframework.integration.ip.tcp.connection.TcpNioConnection.ChannelInputStream; import org.springframework.integration.ip.tcp.serializer.ByteArrayCrLfSerializer; import org.springframework.integration.support.MessageBuilder; import org.springframework.integration.test.util.SocketUtils; @@ -337,6 +339,109 @@ public class TcpNioConnectionTests { assertTrue(messageLatch.await(10, TimeUnit.SECONDS)); } + @Test + public void testByteArrayRead() throws Exception { + SocketChannel socketChannel = mock(SocketChannel.class); + Socket socket = mock(Socket.class); + when(socketChannel.socket()).thenReturn(socket); + TcpNioConnection connection = new TcpNioConnection(socketChannel, false, false, null, null); + TcpNioConnection.ChannelInputStream stream = (ChannelInputStream) new DirectFieldAccessor(connection) + .getPropertyValue("channelInputStream"); + stream.write("foo".getBytes(), 3); + byte[] out = new byte[2]; + int n = stream.read(out); + assertEquals(2, n); + assertEquals("fo", new String(out)); + out = new byte[2]; + n = stream.read(out); + assertEquals(1, n); + assertEquals("o\u0000", new String(out)); + } + + @Test + public void testByteArrayReadMulti() throws Exception { + SocketChannel socketChannel = mock(SocketChannel.class); + Socket socket = mock(Socket.class); + when(socketChannel.socket()).thenReturn(socket); + TcpNioConnection connection = new TcpNioConnection(socketChannel, false, false, null, null); + TcpNioConnection.ChannelInputStream stream = (ChannelInputStream) new DirectFieldAccessor(connection) + .getPropertyValue("channelInputStream"); + stream.write("foo".getBytes(), 3); + stream.write("bar".getBytes(), 3); + byte[] out = new byte[6]; + int n = stream.read(out); + assertEquals(6, n); + assertEquals("foobar", new String(out)); + } + + @Test + public void testByteArrayReadWithOffset() throws Exception { + SocketChannel socketChannel = mock(SocketChannel.class); + Socket socket = mock(Socket.class); + when(socketChannel.socket()).thenReturn(socket); + TcpNioConnection connection = new TcpNioConnection(socketChannel, false, false, null, null); + TcpNioConnection.ChannelInputStream stream = (ChannelInputStream) new DirectFieldAccessor(connection) + .getPropertyValue("channelInputStream"); + stream.write("foo".getBytes(), 3); + byte[] out = new byte[5]; + int n = stream.read(out, 1, 4); + assertEquals(3, n); + assertEquals("\u0000foo\u0000", new String(out)); + } + + @Test + public void testByteArrayReadWithBadArgs() throws Exception { + SocketChannel socketChannel = mock(SocketChannel.class); + Socket socket = mock(Socket.class); + when(socketChannel.socket()).thenReturn(socket); + TcpNioConnection connection = new TcpNioConnection(socketChannel, false, false, null, null); + TcpNioConnection.ChannelInputStream stream = (ChannelInputStream) new DirectFieldAccessor(connection) + .getPropertyValue("channelInputStream"); + stream.write("foo".getBytes(), 3); + byte[] out = new byte[5]; + try { + stream.read(out, 1, 5); + fail("Expected IndexOutOfBoundsException"); + } + catch (IndexOutOfBoundsException e) {} + try { + stream.read(null, 1, 5); + fail("Expected IllegalArgumentException"); + } + catch (IllegalArgumentException e) {} + assertEquals(0, stream.read(out, 0, 0)); + assertEquals(3, stream.read(out)); + } + + @Test + public void testByteArrayBlocksForZeroRead() throws Exception { + SocketChannel socketChannel = mock(SocketChannel.class); + Socket socket = mock(Socket.class); + when(socketChannel.socket()).thenReturn(socket); + TcpNioConnection connection = new TcpNioConnection(socketChannel, false, false, null, null); + final TcpNioConnection.ChannelInputStream stream = (ChannelInputStream) new DirectFieldAccessor(connection) + .getPropertyValue("channelInputStream"); + final CountDownLatch latch = new CountDownLatch(1); + final byte[] out = new byte[4]; + ExecutorService exec = Executors.newSingleThreadExecutor(); + exec.execute(new Runnable(){ + public void run() { + try { + stream.read(out); + } + catch (IOException e) { + e.printStackTrace(); + } + latch.countDown(); + } + }); + Thread.sleep(1000); + assertEquals(0x00, out[0]); + stream.write("foo".getBytes(), 3); + assertTrue(latch.await(10, TimeUnit.SECONDS)); + assertEquals("foo\u0000", new String(out)); + } + private void readFully(InputStream is, byte[] buff) throws IOException { for (int i = 0; i < buff.length; i++) { buff[i] = (byte) is.read();