diff --git a/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/connection/TcpNioConnectionReadTests.java b/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/connection/TcpNioConnectionReadTests.java index e85f9142e3..295b8ff625 100644 --- a/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/connection/TcpNioConnectionReadTests.java +++ b/spring-integration-ip/src/test/java/org/springframework/integration/ip/tcp/connection/TcpNioConnectionReadTests.java @@ -16,7 +16,10 @@ package org.springframework.integration.ip.tcp.connection; +import static org.hamcrest.Matchers.anyOf; +import static org.hamcrest.Matchers.containsString; import static org.junit.Assert.assertEquals; +import static org.junit.Assert.assertThat; import static org.junit.Assert.assertTrue; import static org.junit.Assert.fail; import static org.mockito.Mockito.mock; @@ -27,6 +30,7 @@ import java.util.List; import java.util.concurrent.CountDownLatch; import java.util.concurrent.Semaphore; import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicReference; import javax.net.SocketFactory; @@ -40,9 +44,12 @@ import org.springframework.integration.ip.tcp.serializer.ByteArrayStxEtxSerializ import org.springframework.integration.ip.util.SocketTestUtils; import org.springframework.integration.ip.util.TestingUtilities; import org.springframework.messaging.Message; +import org.springframework.messaging.support.ErrorMessage; /** * @author Gary Russell + * @author Artem Bilan + * * @since 2.0 */ public class TcpNioConnectionReadTests { @@ -83,7 +90,6 @@ public class TcpNioConnectionReadTests { semaphore.release(); return false; } - }); // Fire up the sender. @@ -96,13 +102,12 @@ public class TcpNioConnectionReadTests { assertEquals("Data", SocketTestUtils.TEST_STRING + SocketTestUtils.TEST_STRING, new String((byte[]) responses.get(0).getPayload())); assertEquals("Data", SocketTestUtils.TEST_STRING + SocketTestUtils.TEST_STRING, - new String((byte[]) responses.get(1).getPayload())); + new String((byte[]) responses.get(1).getPayload())); scf.stop(); done.countDown(); } - @SuppressWarnings("unchecked") @Test public void testFragmented() throws Exception { @@ -115,7 +120,7 @@ public class TcpNioConnectionReadTests { public boolean onMessage(Message message) { responses.add(message); try { - Thread.sleep(1000); + Thread.sleep(10); } catch (InterruptedException e) { Thread.currentThread().interrupt(); @@ -123,7 +128,6 @@ public class TcpNioConnectionReadTests { semaphore.release(); return false; } - }); int howMany = 2; @@ -134,7 +138,7 @@ public class TcpNioConnectionReadTests { assertEquals("Expected", howMany, responses.size()); for (int i = 0; i < howMany; i++) { assertEquals("Data", "xx", - new String(((Message) responses.get(0)).getPayload())); + new String(((Message) responses.get(0)).getPayload())); } scf.stop(); done.countDown(); @@ -154,7 +158,6 @@ public class TcpNioConnectionReadTests { semaphore.release(); return false; } - }); // Fire up the sender. @@ -165,9 +168,9 @@ public class TcpNioConnectionReadTests { assertTrue(semaphore.tryAcquire(1, 10000, TimeUnit.MILLISECONDS)); assertEquals("Did not receive data", 2, responses.size()); assertEquals("Data", SocketTestUtils.TEST_STRING + SocketTestUtils.TEST_STRING, - new String(((Message) responses.get(0)).getPayload())); + new String(((Message) responses.get(0)).getPayload())); assertEquals("Data", SocketTestUtils.TEST_STRING + SocketTestUtils.TEST_STRING, - new String(((Message) responses.get(1)).getPayload())); + new String(((Message) responses.get(1)).getPayload())); scf.stop(); done.countDown(); } @@ -186,7 +189,6 @@ public class TcpNioConnectionReadTests { semaphore.release(); return false; } - }); // Fire up the sender. @@ -197,9 +199,9 @@ public class TcpNioConnectionReadTests { assertTrue(semaphore.tryAcquire(1, 10000, TimeUnit.MILLISECONDS)); assertEquals("Did not receive data", 2, responses.size()); assertEquals("Data", SocketTestUtils.TEST_STRING + SocketTestUtils.TEST_STRING, - new String(((Message) responses.get(0)).getPayload())); + new String(((Message) responses.get(0)).getPayload())); assertEquals("Data", SocketTestUtils.TEST_STRING + SocketTestUtils.TEST_STRING, - new String(((Message) responses.get(1)).getPayload())); + new String(((Message) responses.get(1)).getPayload())); scf.stop(); done.countDown(); } @@ -210,14 +212,20 @@ public class TcpNioConnectionReadTests { final Semaphore semaphore = new Semaphore(0); final List added = new ArrayList(); final List removed = new ArrayList(); + + final CountDownLatch errorMessageLetch = new CountDownLatch(1); + final AtomicReference errorMessageRef = new AtomicReference(); + AbstractServerConnectionFactory scf = getConnectionFactory(serializer, new TcpListener() { @Override public boolean onMessage(Message message) { - semaphore.release(); + if (message instanceof ErrorMessage) { + errorMessageRef.set(((ErrorMessage) message).getPayload()); + errorMessageLetch.countDown(); + } return false; } - }, new TcpSender() { @Override @@ -239,6 +247,13 @@ public class TcpNioConnectionReadTests { CountDownLatch done = SocketTestUtils.testSendLengthOverflow(scf.getPort()); whileOpen(semaphore, added); assertEquals(1, added.size()); + + assertTrue(errorMessageLetch.await(10, TimeUnit.SECONDS)); + + assertThat(errorMessageRef.get().getMessage(), + anyOf(containsString("Message length 2147483647 exceeds max message length: 2048"), + containsString("Connection is closed"))); + assertTrue(semaphore.tryAcquire(10000, TimeUnit.MILLISECONDS)); assertTrue(removed.size() > 0); scf.stop(); @@ -252,14 +267,20 @@ public class TcpNioConnectionReadTests { final Semaphore semaphore = new Semaphore(0); final List added = new ArrayList(); final List removed = new ArrayList(); + + final CountDownLatch errorMessageLetch = new CountDownLatch(1); + final AtomicReference errorMessageRef = new AtomicReference(); + AbstractServerConnectionFactory scf = getConnectionFactory(serializer, new TcpListener() { @Override public boolean onMessage(Message message) { - semaphore.release(); + if (message instanceof ErrorMessage) { + errorMessageRef.set(((ErrorMessage) message).getPayload()); + errorMessageLetch.countDown(); + } return false; } - }, new TcpSender() { @Override @@ -281,6 +302,13 @@ public class TcpNioConnectionReadTests { CountDownLatch done = SocketTestUtils.testSendStxEtxOverflow(scf.getPort()); whileOpen(semaphore, added); assertEquals(1, added.size()); + + assertTrue(errorMessageLetch.await(10, TimeUnit.SECONDS)); + + assertThat(errorMessageRef.get().getMessage(), + anyOf(containsString("Connection is closed"), + containsString("ETX not found before max message length: 1024"))); + assertTrue(semaphore.tryAcquire(10000, TimeUnit.MILLISECONDS)); assertTrue(removed.size() > 0); scf.stop(); @@ -294,14 +322,20 @@ public class TcpNioConnectionReadTests { final Semaphore semaphore = new Semaphore(0); final List added = new ArrayList(); final List removed = new ArrayList(); + + final CountDownLatch errorMessageLetch = new CountDownLatch(1); + final AtomicReference errorMessageRef = new AtomicReference(); + AbstractServerConnectionFactory scf = getConnectionFactory(serializer, new TcpListener() { @Override public boolean onMessage(Message message) { - semaphore.release(); + if (message instanceof ErrorMessage) { + errorMessageRef.set(((ErrorMessage) message).getPayload()); + errorMessageLetch.countDown(); + } return false; } - }, new TcpSender() { @Override @@ -323,6 +357,13 @@ public class TcpNioConnectionReadTests { CountDownLatch done = SocketTestUtils.testSendCrLfOverflow(scf.getPort()); whileOpen(semaphore, added); assertEquals(1, added.size()); + + assertTrue(errorMessageLetch.await(10, TimeUnit.SECONDS)); + + assertThat(errorMessageRef.get().getMessage(), + anyOf(containsString("Connection is closed"), + containsString("CRLF not found before max message length: 1024"))); + assertTrue(semaphore.tryAcquire(10000, TimeUnit.MILLISECONDS)); assertTrue(removed.size() > 0); scf.stop(); @@ -331,7 +372,6 @@ public class TcpNioConnectionReadTests { /** * Tests socket closure when no data received. - * * @throws Exception */ @Test @@ -341,14 +381,20 @@ public class TcpNioConnectionReadTests { final Semaphore semaphore = new Semaphore(0); final List added = new ArrayList(); final List removed = new ArrayList(); + + final CountDownLatch errorMessageLetch = new CountDownLatch(1); + final AtomicReference errorMessageRef = new AtomicReference(); + AbstractServerConnectionFactory scf = getConnectionFactory(serializer, new TcpListener() { @Override public boolean onMessage(Message message) { - semaphore.release(); + if (message instanceof ErrorMessage) { + errorMessageRef.set(((ErrorMessage) message).getPayload()); + errorMessageLetch.countDown(); + } return false; } - }, new TcpSender() { @Override @@ -368,6 +414,12 @@ public class TcpNioConnectionReadTests { socket.close(); whileOpen(semaphore, added); assertEquals(1, added.size()); + + assertTrue(errorMessageLetch.await(10, TimeUnit.SECONDS)); + + assertThat(errorMessageRef.get().getMessage(), + anyOf(containsString("Connection is closed"), containsString("Stream closed after 2 of 3"))); + assertTrue(semaphore.tryAcquire(10000, TimeUnit.MILLISECONDS)); assertTrue(removed.size() > 0); scf.stop(); @@ -375,7 +427,6 @@ public class TcpNioConnectionReadTests { /** * Tests socket closure when no data received. - * * @throws Exception */ @Test @@ -385,14 +436,20 @@ public class TcpNioConnectionReadTests { final Semaphore semaphore = new Semaphore(0); final List added = new ArrayList(); final List removed = new ArrayList(); + + final CountDownLatch errorMessageLetch = new CountDownLatch(1); + final AtomicReference errorMessageRef = new AtomicReference(); + AbstractServerConnectionFactory scf = getConnectionFactory(serializer, new TcpListener() { @Override public boolean onMessage(Message message) { - semaphore.release(); + if (message instanceof ErrorMessage) { + errorMessageRef.set(((ErrorMessage) message).getPayload()); + errorMessageLetch.countDown(); + } return false; } - }, new TcpSender() { @Override @@ -413,6 +470,12 @@ public class TcpNioConnectionReadTests { socket.close(); whileOpen(semaphore, added); assertEquals(1, added.size()); + + assertTrue(errorMessageLetch.await(10, TimeUnit.SECONDS)); + + assertThat(errorMessageRef.get().getMessage(), + anyOf(containsString("Connection is closed"), containsString("Socket closed during message assembly"))); + assertTrue(semaphore.tryAcquire(10000, TimeUnit.MILLISECONDS)); assertTrue(removed.size() > 0); scf.stop(); @@ -420,7 +483,6 @@ public class TcpNioConnectionReadTests { /** * Tests socket closure when mid-message - * * @throws Exception */ @Test @@ -431,10 +493,8 @@ public class TcpNioConnectionReadTests { /** * Tests socket closure when mid-message - * * @throws Exception */ - @Test public void testCloseCleanupStxEtx() throws Exception { ByteArrayCrLfSerializer serializer = new ByteArrayCrLfSerializer(); @@ -443,10 +503,8 @@ public class TcpNioConnectionReadTests { /** * Tests socket closure when mid-message - * * @throws Exception */ - @Test public void testCloseCleanupLengthHeader() throws Exception { ByteArrayLengthHeaderSerializer serializer = new ByteArrayLengthHeaderSerializer(); @@ -455,33 +513,51 @@ public class TcpNioConnectionReadTests { private void testClosureMidMessageGuts(AbstractByteArraySerializer serializer, String shortMessage) throws Exception { - final List> responses = new ArrayList>(); final Semaphore semaphore = new Semaphore(0); final List added = new ArrayList(); final List removed = new ArrayList(); + + final CountDownLatch errorMessageLetch = new CountDownLatch(1); + final AtomicReference errorMessageRef = new AtomicReference(); + AbstractServerConnectionFactory scf = getConnectionFactory(serializer, new TcpListener() { + @Override public boolean onMessage(Message message) { - responses.add(message); + if (message instanceof ErrorMessage) { + errorMessageRef.set(((ErrorMessage) message).getPayload()); + errorMessageLetch.countDown(); + } return false; } }, new TcpSender() { + @Override public void addNewConnection(TcpConnection connection) { added.add(connection); semaphore.release(); } + @Override public void removeDeadConnection(TcpConnection connection) { removed.add(connection); semaphore.release(); } + }); Socket socket = SocketFactory.getDefault().createSocket("localhost", scf.getPort()); socket.getOutputStream().write(shortMessage.getBytes()); socket.close(); whileOpen(semaphore, added); assertEquals(1, added.size()); + + assertTrue(errorMessageLetch.await(10, TimeUnit.SECONDS)); + + assertThat(errorMessageRef.get().getMessage(), + anyOf(containsString("Connection is closed"), + containsString("Socket closed during message assembly"), + containsString("Stream closed after 2 of 3"))); + assertTrue(semaphore.tryAcquire(10000, TimeUnit.MILLISECONDS)); assertTrue(removed.size() > 0); scf.stop();