diff --git a/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/connection/TcpNetConnection.java b/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/connection/TcpNetConnection.java index 8a5fa0c041..1c4e6f76bf 100644 --- a/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/connection/TcpNetConnection.java +++ b/spring-integration-ip/src/main/java/org/springframework/integration/ip/tcp/connection/TcpNetConnection.java @@ -72,7 +72,13 @@ public class TcpNetConnection extends AbstractTcpConnection { public synchronized void send(Message message) throws Exception { Object object = this.getMapper().fromMessage(message); this.lastSend = System.currentTimeMillis(); - ((Serializer) this.getSerializer()).serialize(object, this.socket.getOutputStream()); + try { + ((Serializer) this.getSerializer()).serialize(object, this.socket.getOutputStream()); + } + catch (Exception e) { + this.closeConnection(true); + throw e; + } this.afterSend(message); } 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 2437589720..9f8e9a013c 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 @@ -118,10 +118,16 @@ public class TcpNioConnection extends AbstractTcpConnection { @SuppressWarnings("unchecked") public void send(Message message) throws Exception { - synchronized(this.getMapper()) { + synchronized(this.socketChannel) { Object object = this.getMapper().fromMessage(message); this.lastSend = System.currentTimeMillis(); - ((Serializer) this.getSerializer()).serialize(object, this.getChannelOutputStream()); + try { + ((Serializer) this.getSerializer()).serialize(object, this.getChannelOutputStream()); + } + catch (Exception e) { + this.closeConnection(true); + throw e; + } this.afterSend(message); } } diff --git a/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/connection/CachingClientConnectionFactoryTests.java b/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/connection/CachingClientConnectionFactoryTests.java index 28edb5ef00..414c9ae135 100644 --- a/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/connection/CachingClientConnectionFactoryTests.java +++ b/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/connection/CachingClientConnectionFactoryTests.java @@ -20,12 +20,19 @@ import static org.junit.Assert.assertFalse; import static org.junit.Assert.assertNotNull; import static org.junit.Assert.assertNotSame; import static org.junit.Assert.assertSame; +import static org.junit.Assert.assertTrue; +import static org.junit.Assert.fail; import static org.mockito.Mockito.doAnswer; import static org.mockito.Mockito.doThrow; import static org.mockito.Mockito.mock; import static org.mockito.Mockito.verify; import static org.mockito.Mockito.when; +import java.io.IOException; +import java.io.OutputStream; +import java.net.Socket; +import java.nio.ByteBuffer; +import java.nio.channels.SocketChannel; import java.util.ArrayList; import java.util.List; import java.util.concurrent.Executors; @@ -34,15 +41,18 @@ import java.util.concurrent.atomic.AtomicBoolean; import org.junit.Test; import org.junit.runner.RunWith; +import org.mockito.Mockito; import org.mockito.invocation.InvocationOnMock; import org.mockito.stubbing.Answer; +import org.springframework.beans.DirectFieldAccessor; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.integration.Message; import org.springframework.integration.MessagingException; import org.springframework.integration.core.PollableChannel; import org.springframework.integration.core.SubscribableChannel; import org.springframework.integration.ip.IpHeaders; +import org.springframework.integration.ip.tcp.serializer.ByteArrayCrLfSerializer; import org.springframework.integration.ip.util.TestingUtilities; import org.springframework.integration.message.GenericMessage; import org.springframework.integration.support.MessageBuilder; @@ -281,6 +291,74 @@ public class CachingClientConnectionFactoryTests { verify(mockConn2).close(); } + @Test + public void testExceptionOnSendNet() throws Exception { + AbstractTcpConnection conn1 = mockedTcpNetConnection(); + AbstractTcpConnection conn2 = mockedTcpNetConnection(); + + CachingClientConnectionFactory cccf = createCCCFWith2Connections(conn1, conn2); + doTestCloseOnSendError(conn1, conn2, cccf); + } + + @Test + public void testExceptionOnSendNio() throws Exception { + AbstractTcpConnection conn1 = mockedTcpNioConnection(); + AbstractTcpConnection conn2 = mockedTcpNioConnection(); + + CachingClientConnectionFactory cccf = createCCCFWith2Connections(conn1, conn2); + doTestCloseOnSendError(conn1, conn2, cccf); + } + + private void doTestCloseOnSendError(TcpConnection conn1, TcpConnection conn2, + CachingClientConnectionFactory cccf) throws Exception { + TcpConnection cached1 = cccf.getConnection(); + try { + cached1.send(new GenericMessage("foo")); + fail("Expected IOException"); + } + catch (IOException e) { + assertEquals("Foo", e.getMessage()); + } + // Before INT-3163 this failed with a timeout - connection not returned to pool after failure on send() + TcpConnection cached2 = cccf.getConnection(); + assertTrue(cached1.getConnectionId().contains(conn1.getConnectionId())); + assertTrue(cached2.getConnectionId().contains(conn2.getConnectionId())); + } + + private CachingClientConnectionFactory createCCCFWith2Connections(AbstractTcpConnection conn1, AbstractTcpConnection conn2) + throws Exception { + AbstractClientConnectionFactory factory = mock(AbstractClientConnectionFactory.class); + when(factory.isRunning()).thenReturn(true); + when(factory.getConnection()).thenReturn(conn1, conn2); + CachingClientConnectionFactory cccf = new CachingClientConnectionFactory(factory, 1); + cccf.setConnectionWaitTimeout(1); + cccf.start(); + return cccf; + } + + private AbstractTcpConnection mockedTcpNetConnection() throws IOException { + Socket socket = mock(Socket.class); + when(socket.isClosed()).thenReturn(true); // closed when next retrieved + OutputStream stream = mock(OutputStream.class); + doThrow(new IOException("Foo")).when(stream).write(Mockito.any(byte[].class)); + when(socket.getOutputStream()).thenReturn(stream); + TcpNetConnection conn = new TcpNetConnection(socket, false, false); + conn.setMapper(new TcpMessageMapper()); + conn.setSerializer(new ByteArrayCrLfSerializer()); + return conn; + } + + private AbstractTcpConnection mockedTcpNioConnection() throws Exception { + SocketChannel socketChannel = mock(SocketChannel.class); + new DirectFieldAccessor(socketChannel).setPropertyValue("open", false); + doThrow(new IOException("Foo")).when(socketChannel).write(Mockito.any(ByteBuffer.class)); + when(socketChannel.socket()).thenReturn(mock(Socket.class)); + TcpNioConnection conn = new TcpNioConnection(socketChannel, false, false); + conn.setMapper(new TcpMessageMapper()); + conn.setSerializer(new ByteArrayCrLfSerializer()); + return conn; + } + private TcpConnection makeMockConnection(String name) { return makeMockConnection(name, false); }