This is an automated email from the ASF dual-hosted git repository.
rzo1 pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/tomee.git
The following commit(s) were added to refs/heads/main by this push:
new 88f709bc99 restore mp-jwt security state when the filter chain throws
(#3058)
88f709bc99 is described below
commit 88f709bc993c89cb45aa2fe6e7d0ddd5a7a6ee2f
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.
---
.../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();
+ }
+ }
+}