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();
+        }
+    }
+}

Reply via email to