diff --git a/spring-boot/src/main/java/org/springframework/boot/context/web/ErrorPageFilter.java b/spring-boot/src/main/java/org/springframework/boot/context/web/ErrorPageFilter.java index f6af4d60ce..c772c3ba2a 100644 --- a/spring-boot/src/main/java/org/springframework/boot/context/web/ErrorPageFilter.java +++ b/spring-boot/src/main/java/org/springframework/boot/context/web/ErrorPageFilter.java @@ -39,6 +39,7 @@ import org.springframework.core.Ordered; import org.springframework.core.annotation.Order; import org.springframework.stereotype.Component; import org.springframework.web.filter.OncePerRequestFilter; +import org.springframework.web.util.NestedServletException; /** * A special {@link AbstractConfigurableEmbeddedServletContainer} for non-embedded @@ -69,6 +70,8 @@ public class ErrorPageFilter extends AbstractConfigurableEmbeddedServletContaine private static final String ERROR_MESSAGE = "javax.servlet.error.message"; + public static final String ERROR_REQUEST_URI = "javax.servlet.error.request_uri"; + private static final String ERROR_STATUS_CODE = "javax.servlet.error.status_code"; private String global; @@ -121,7 +124,11 @@ public class ErrorPageFilter extends AbstractConfigurableEmbeddedServletContaine } } catch (Throwable ex) { - handleException(request, response, wrapped, ex); + Throwable exceptionToHandle = ex; + if (ex instanceof NestedServletException) { + exceptionToHandle = ((NestedServletException) ex).getRootCause(); + } + handleException(request, response, wrapped, exceptionToHandle); response.flushBuffer(); } } @@ -225,9 +232,10 @@ public class ErrorPageFilter extends AbstractConfigurableEmbeddedServletContaine return this.global; } - private void setErrorAttributes(ServletRequest request, int status, String message) { + private void setErrorAttributes(HttpServletRequest request, int status, String message) { request.setAttribute(ERROR_STATUS_CODE, status); request.setAttribute(ERROR_MESSAGE, message); + request.setAttribute(ERROR_REQUEST_URI, request.getRequestURI()); } private void rethrow(Throwable ex) throws IOException, ServletException { diff --git a/spring-boot/src/test/java/org/springframework/boot/context/web/ErrorPageFilterTests.java b/spring-boot/src/test/java/org/springframework/boot/context/web/ErrorPageFilterTests.java index 2e5d18e485..393889fa5e 100644 --- a/spring-boot/src/test/java/org/springframework/boot/context/web/ErrorPageFilterTests.java +++ b/spring-boot/src/test/java/org/springframework/boot/context/web/ErrorPageFilterTests.java @@ -38,6 +38,7 @@ import org.springframework.web.context.request.async.DeferredResult; import org.springframework.web.context.request.async.StandardServletAsyncWebRequest; import org.springframework.web.context.request.async.WebAsyncManager; import org.springframework.web.context.request.async.WebAsyncUtils; +import org.springframework.web.util.NestedServletException; import static org.hamcrest.Matchers.containsString; import static org.hamcrest.Matchers.equalTo; @@ -62,7 +63,8 @@ public class ErrorPageFilterTests { private ErrorPageFilter filter = new ErrorPageFilter(); - private MockHttpServletRequest request = new MockHttpServletRequest(); + private MockHttpServletRequest request = new MockHttpServletRequest("GET", + "/test/path"); private MockHttpServletResponse response = new MockHttpServletResponse(); @@ -199,6 +201,9 @@ public class ErrorPageFilterTests { equalTo((Object) 400)); assertThat(this.request.getAttribute(RequestDispatcher.ERROR_MESSAGE), equalTo((Object) "BAD")); + assertThat(this.request.getAttribute(RequestDispatcher.ERROR_REQUEST_URI), + equalTo((Object) "/test/path")); + assertTrue(this.response.isCommitted()); assertThat(this.response.getForwardedUrl(), equalTo("/error")); } @@ -221,6 +226,8 @@ public class ErrorPageFilterTests { equalTo((Object) 400)); assertThat(this.request.getAttribute(RequestDispatcher.ERROR_MESSAGE), equalTo((Object) "BAD")); + assertThat(this.request.getAttribute(RequestDispatcher.ERROR_REQUEST_URI), + equalTo((Object) "/test/path")); assertTrue(this.response.isCommitted()); assertThat(this.response.getForwardedUrl(), equalTo("/400")); } @@ -264,6 +271,8 @@ public class ErrorPageFilterTests { equalTo((Object) "BAD")); assertThat(this.request.getAttribute(RequestDispatcher.ERROR_EXCEPTION_TYPE), equalTo((Object) RuntimeException.class.getName())); + assertThat(this.request.getAttribute(RequestDispatcher.ERROR_REQUEST_URI), + equalTo((Object) "/test/path")); assertTrue(this.response.isCommitted()); assertThat(this.response.getForwardedUrl(), equalTo("/500")); } @@ -319,6 +328,8 @@ public class ErrorPageFilterTests { equalTo((Object) "BAD")); assertThat(this.request.getAttribute(RequestDispatcher.ERROR_EXCEPTION_TYPE), equalTo((Object) IllegalStateException.class.getName())); + assertThat(this.request.getAttribute(RequestDispatcher.ERROR_REQUEST_URI), + equalTo((Object) "/test/path")); assertTrue(this.response.isCommitted()); } @@ -465,6 +476,32 @@ public class ErrorPageFilterTests { assertThat(this.output.toString(), containsString("request [/test/alpha]")); } + @Test + public void nestedServletExceptionIsUnwrapped() throws Exception { + this.filter.addErrorPages(new ErrorPage(RuntimeException.class, "/500")); + this.chain = new MockFilterChain() { + @Override + public void doFilter(ServletRequest request, ServletResponse response) + throws IOException, ServletException { + super.doFilter(request, response); + throw new NestedServletException("Wrapper", new RuntimeException("BAD")); + } + }; + this.filter.doFilter(this.request, this.response, this.chain); + assertThat(((HttpServletResponseWrapper) this.chain.getResponse()).getStatus(), + equalTo(500)); + assertThat(this.request.getAttribute(RequestDispatcher.ERROR_STATUS_CODE), + equalTo((Object) 500)); + assertThat(this.request.getAttribute(RequestDispatcher.ERROR_MESSAGE), + equalTo((Object) "BAD")); + assertThat(this.request.getAttribute(RequestDispatcher.ERROR_EXCEPTION_TYPE), + equalTo((Object) RuntimeException.class.getName())); + assertThat(this.request.getAttribute(RequestDispatcher.ERROR_REQUEST_URI), + equalTo((Object) "/test/path")); + assertTrue(this.response.isCommitted()); + assertThat(this.response.getForwardedUrl(), equalTo("/500")); + } + private void setUpAsyncDispatch() throws Exception { this.request.setAsyncSupported(true); this.request.setAsyncStarted(true);