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.
This commit is contained in:
Gary Russell
2013-03-11 22:34:14 -04:00
committed by Mark Fisher
parent 3eab287e95
commit 474a00509b
2 changed files with 132 additions and 0 deletions

View File

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

View File

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