diff --git a/spring-boot-project/spring-boot/src/test/java/org/springframework/boot/web/embedded/tomcat/TomcatServletWebServerFactoryTests.java b/spring-boot-project/spring-boot/src/test/java/org/springframework/boot/web/embedded/tomcat/TomcatServletWebServerFactoryTests.java index bdb520520f..aec1b7fa00 100644 --- a/spring-boot-project/spring-boot/src/test/java/org/springframework/boot/web/embedded/tomcat/TomcatServletWebServerFactoryTests.java +++ b/spring-boot-project/spring-boot/src/test/java/org/springframework/boot/web/embedded/tomcat/TomcatServletWebServerFactoryTests.java @@ -66,6 +66,7 @@ import org.apache.jasper.servlet.JspServlet; import org.apache.tomcat.JarScanFilter; import org.apache.tomcat.JarScanType; import org.assertj.core.api.ThrowableAssert.ThrowingCallable; +import org.awaitility.Awaitility; import org.junit.jupiter.api.AfterEach; import org.junit.jupiter.api.Test; import org.mockito.ArgumentCaptor; @@ -599,17 +600,21 @@ class TomcatServletWebServerFactoryTests extends AbstractServletWebServerFactory assertThat(keepAliveRequest.get()).isInstanceOf(HttpResponse.class); Future request = initiateGetRequest(port, "/blocking"); blockingServlet.awaitQueue(); + blockingServlet.setBlocking(false); this.webServer.shutDownGracefully((result) -> { }); - Future idleConnectionRequest = initiateGetRequest(httpClient, port, "/blocking"); - blockingServlet.admitOne(); - Object response = request.get(); - assertThat(response).isInstanceOf(HttpResponse.class); - Object idleConnectionRequestResult = idleConnectionRequest.get(); + Object idleConnectionRequestResult = Awaitility.await().until(() -> { + Future idleConnectionRequest = initiateGetRequest(httpClient, port, "/blocking"); + Object result = idleConnectionRequest.get(); + return result; + }, (result) -> result instanceof Exception); assertThat(idleConnectionRequestResult).isInstanceOfAny(SocketException.class, NoHttpResponseException.class); if (idleConnectionRequestResult instanceof SocketException) { assertThat((SocketException) idleConnectionRequestResult).hasMessage("Connection reset"); } + blockingServlet.admitOne(); + Object response = request.get(); + assertThat(response).isInstanceOf(HttpResponse.class); this.webServer.stop(); } diff --git a/spring-boot-project/spring-boot/src/test/java/org/springframework/boot/web/servlet/server/AbstractServletWebServerFactoryTests.java b/spring-boot-project/spring-boot/src/test/java/org/springframework/boot/web/servlet/server/AbstractServletWebServerFactoryTests.java index c51f1f0f19..655617df78 100644 --- a/spring-boot-project/spring-boot/src/test/java/org/springframework/boot/web/servlet/server/AbstractServletWebServerFactoryTests.java +++ b/spring-boot-project/spring-boot/src/test/java/org/springframework/boot/web/servlet/server/AbstractServletWebServerFactoryTests.java @@ -1445,22 +1445,26 @@ public abstract class AbstractServletWebServerFactoryTests { private final BlockingQueue barriers = new ArrayBlockingQueue<>(10); + protected volatile boolean blocking = true; + public BlockingServlet() { } @Override protected void doGet(HttpServletRequest req, HttpServletResponse resp) throws ServletException, IOException { - CyclicBarrier barrier = new CyclicBarrier(2); - this.barriers.add(barrier); - try { - barrier.await(); - } - catch (InterruptedException ex) { - Thread.currentThread().interrupt(); - } - catch (BrokenBarrierException ex) { - throw new ServletException(ex); + if (this.blocking) { + CyclicBarrier barrier = new CyclicBarrier(2); + this.barriers.add(barrier); + try { + barrier.await(); + } + catch (InterruptedException ex) { + Thread.currentThread().interrupt(); + } + catch (BrokenBarrierException ex) { + throw new ServletException(ex); + } } } @@ -1491,6 +1495,10 @@ public abstract class AbstractServletWebServerFactoryTests { } } + public void setBlocking(boolean blocking) { + this.blocking = blocking; + } + } static class BlockingAsyncServlet extends HttpServlet {