This is an automated email from the ASF dual-hosted git repository. rzo1 pushed a commit to branch tomee-10.x in repository https://gitbox.apache.org/repos/asf/tomee.git
commit 7a7f32455223c2e46efad5f6ae90249ff74a65db Author: Richard Zowalla <[email protected]> AuthorDate: Wed Oct 7 17:34:32 2026 +0200 restore mp-jwt security state when the filter chain throws (#3058) Move exitWebApp into a finally block so the caller identity and run-as subject are not left on the pooled thread when the chain throws. (cherry picked from commit 88f709bc993c89cb45aa2fe6e7d0ddd5a7a6ee2f) --- .../apache/tomee/microprofile/jwt/MPJWTFilter.java | 31 +++- .../jwt/MPJWTFilterExitWebAppTest.java | 186 +++++++++++++++++++++ 2 files changed, 209 insertions(+), 8 deletions(-) diff --git a/mp-jwt/src/main/java/org/apache/tomee/microprofile/jwt/MPJWTFilter.java b/mp-jwt/src/main/java/org/apache/tomee/microprofile/jwt/MPJWTFilter.java index 95a739b0e7..170bcabcef 100644 --- a/mp-jwt/src/main/java/org/apache/tomee/microprofile/jwt/MPJWTFilter.java +++ b/mp-jwt/src/main/java/org/apache/tomee/microprofile/jwt/MPJWTFilter.java @@ -80,6 +80,8 @@ public class MPJWTFilter implements Filter { private static final Logger VALIDATION = Logger.getInstance(JWTLogCategories.VALIDATION, MPJWTFilter.class); + static final String PRE_LOGIN_STATE = "MP_JWT_PRE_LOGIN_STATE"; + @Override public void init(final FilterConfig filterConfig) throws ServletException { } @@ -98,13 +100,6 @@ public class MPJWTFilter implements Filter { try { final MPJWTServletRequestWrapper wrappedRequest = new MPJWTServletRequestWrapper(httpServletRequest, authContextInfo.get()); chain.doFilter(wrappedRequest, response); - - Object state = request.getAttribute("MP_JWT_PRE_LOGIN_STATE"); - final SecurityService securityService = SystemInstance.get().getComponent(SecurityService.class); - if (TomcatSecurityService.class.isInstance(securityService) && state != null) { - final TomcatSecurityService tomcatSecurityService = TomcatSecurityService.class.cast(securityService); - tomcatSecurityService.exitWebApp(state); - } } catch (final Exception e) { // this is an alternative to the @Provider bellow which requires registration on the fly // or users to add it into their webapp for scanning or into the Application itself @@ -117,6 +112,26 @@ public class MPJWTFilter implements Filter { } else { throw e; } + } finally { + // token validation pushes the caller identity (and the run-as subject, if any) onto + // the current thread; restore it on every exit path so a failing request does not + // leak its security context to the next request served by this pooled thread + exitWebApp(request); + } + } + + private static void exitWebApp(final ServletRequest request) { + final Object state = request.getAttribute(PRE_LOGIN_STATE); + if (state == null) { + return; + } + + // remove it first so the state is never restored twice for the same request + request.removeAttribute(PRE_LOGIN_STATE); + + final SecurityService securityService = SystemInstance.get().getComponent(SecurityService.class); + if (TomcatSecurityService.class.isInstance(securityService)) { + TomcatSecurityService.class.cast(securityService).exitWebApp(state); } } @@ -362,7 +377,7 @@ public class MPJWTFilter implements Filter { final org.apache.catalina.connector.Request req = OpenEJBSecurityListener.requests.get(); Object state = tomcatSecurityService.enterWebApp(req.getWrapper().getRealm(), jsonWebToken, req.getWrapper().getRunAs()); - request.setAttribute("MP_JWT_PRE_LOGIN_STATE", state); + request.setAttribute(PRE_LOGIN_STATE, state); } // TODO Also check if it is an async request and add a listener to close off the state diff --git a/mp-jwt/src/test/java/org/apache/tomee/microprofile/jwt/MPJWTFilterExitWebAppTest.java b/mp-jwt/src/test/java/org/apache/tomee/microprofile/jwt/MPJWTFilterExitWebAppTest.java new file mode 100644 index 0000000000..675372e1cb --- /dev/null +++ b/mp-jwt/src/test/java/org/apache/tomee/microprofile/jwt/MPJWTFilterExitWebAppTest.java @@ -0,0 +1,186 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.apache.tomee.microprofile.jwt; + +import jakarta.enterprise.inject.Instance; +import jakarta.servlet.FilterChain; +import jakarta.servlet.ServletException; +import jakarta.servlet.http.HttpServletRequest; +import jakarta.servlet.http.HttpServletResponse; +import org.apache.catalina.core.StandardServer; +import org.apache.openejb.loader.SystemInstance; +import org.apache.openejb.spi.SecurityService; +import org.apache.tomee.catalina.TomcatSecurityService; +import org.apache.tomee.loader.TomcatHelper; +import org.apache.tomee.microprofile.jwt.config.JWTAuthConfiguration; +import org.junit.After; +import org.junit.Before; +import org.junit.Test; + +import java.lang.reflect.Field; +import java.lang.reflect.Proxy; +import java.util.ArrayList; +import java.util.HashMap; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; + +import static org.junit.Assert.assertEquals; +import static org.junit.Assert.assertNull; +import static org.junit.Assert.assertSame; +import static org.junit.Assert.assertTrue; +import static org.junit.Assert.fail; + +public class MPJWTFilterExitWebAppTest { + + private RecordingSecurityService securityService; + + @Before + public void setUp() { + SystemInstance.reset(); + TomcatHelper.setServer(new StandardServer()); + securityService = new RecordingSecurityService(); + SystemInstance.get().setComponent(SecurityService.class, securityService); + } + + @After + public void tearDown() { + securityService.clearRunAsStack(); + SystemInstance.reset(); + } + + @Test + public void exitWebAppRunsWhenChainThrows() throws Exception { + final Map<String, Object> attributes = new HashMap<>(); + final HttpServletRequest request = request(attributes); + final IllegalStateException failure = new IllegalStateException("boom"); + + final FilterChain chain = (req, res) -> { + enterWebApp(req); + throw failure; + }; + + try { + filter().doFilter(request, response(), chain); + fail("the chain exception must be propagated"); + } catch (final IllegalStateException e) { + assertSame(failure, e); + } + + assertExited(attributes); + } + + @Test + public void exitWebAppRunsWhenChainThrowsCheckedException() throws Exception { + final Map<String, Object> attributes = new HashMap<>(); + final HttpServletRequest request = request(attributes); + + final FilterChain chain = (req, res) -> { + enterWebApp(req); + throw new ServletException("boom"); + }; + + try { + filter().doFilter(request, response(), chain); + fail("the chain exception must be propagated"); + } catch (final ServletException e) { + assertEquals("boom", e.getMessage()); + } + + assertExited(attributes); + } + + @Test + public void exitWebAppRunsOnSuccess() throws Exception { + final Map<String, Object> attributes = new HashMap<>(); + final HttpServletRequest request = request(attributes); + + filter().doFilter(request, response(), (req, res) -> enterWebApp(req)); + + assertExited(attributes); + } + + private void enterWebApp(final jakarta.servlet.ServletRequest req) { + // simulates what ValidateJSonWebToken does once the token has been validated + final Object state = securityService.enterWebApp(null, null, "admin"); + req.setAttribute(MPJWTFilter.PRE_LOGIN_STATE, state); + assertEquals(1, securityService.runAsStackSize()); + } + + private void assertExited(final Map<String, Object> attributes) { + assertEquals(1, securityService.exited.size()); + assertEquals(0, securityService.runAsStackSize()); + assertNull(attributes.get(MPJWTFilter.PRE_LOGIN_STATE)); + } + + private static MPJWTFilter filter() throws Exception { + final JWTAuthConfiguration configuration = new JWTAuthConfiguration(LinkedHashMap::new, "https://server.example.com", + false, new String[0], LinkedHashMap::new, "Authorization", null, null, null, null, 0); + + final Instance<?> instance = (Instance<?>) Proxy.newProxyInstance( + MPJWTFilterExitWebAppTest.class.getClassLoader(), new Class<?>[]{Instance.class}, + (proxy, method, args) -> switch (method.getName()) { + case "isUnsatisfied" -> false; + case "get" -> configuration; + default -> throw new UnsupportedOperationException(method.getName()); + }); + + final MPJWTFilter filter = new MPJWTFilter(); + final Field field = MPJWTFilter.class.getDeclaredField("authContextInfo"); + field.setAccessible(true); + field.set(filter, instance); + return filter; + } + + private static HttpServletRequest request(final Map<String, Object> attributes) { + return (HttpServletRequest) Proxy.newProxyInstance( + MPJWTFilterExitWebAppTest.class.getClassLoader(), new Class<?>[]{HttpServletRequest.class}, + (proxy, method, args) -> switch (method.getName()) { + case "getAttribute" -> attributes.get((String) args[0]); + case "setAttribute" -> attributes.put((String) args[0], args[1]); + case "removeAttribute" -> attributes.remove((String) args[0]); + default -> throw new UnsupportedOperationException(method.getName()); + }); + } + + private static HttpServletResponse response() { + return (HttpServletResponse) Proxy.newProxyInstance( + MPJWTFilterExitWebAppTest.class.getClassLoader(), new Class<?>[]{HttpServletResponse.class}, + (proxy, method, args) -> { + throw new UnsupportedOperationException(method.getName()); + }); + } + + private static class RecordingSecurityService extends TomcatSecurityService { + + private final List<Object> exited = new ArrayList<>(); + + @Override + public void exitWebApp(final Object state) { + exited.add(state); + super.exitWebApp(state); + } + + int runAsStackSize() { + return RUN_AS_STACK.get().size(); + } + + void clearRunAsStack() { + RUN_AS_STACK.remove(); + } + } +}
