diff --git a/app/src/main/java/org/apache/roller/weblogger/ui/core/filters/SaltValidator.java b/app/src/main/java/org/apache/roller/weblogger/ui/core/filters/SaltValidator.java new file mode 100644 index 0000000000..dbe8abd594 --- /dev/null +++ b/app/src/main/java/org/apache/roller/weblogger/ui/core/filters/SaltValidator.java @@ -0,0 +1,91 @@ +/* + * 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. For additional information regarding + * copyright in this work, please see the NOTICE file in the top level + * directory of this distribution. + */ + +package org.apache.roller.weblogger.ui.core.filters; + +import java.util.Locale; +import java.util.Objects; + +import javax.servlet.http.HttpServletRequest; + +import org.apache.roller.weblogger.ui.core.RollerSession; +import org.apache.roller.weblogger.ui.rendering.util.cache.SaltCache; + +/** + * Shared validation for salts submitted by UI forms. + */ +public final class SaltValidator { + + private static final String MULTIPART_FORM_DATA = "multipart/form-data"; + + private SaltValidator() { + } + + /** + * Validates and consumes the salt submitted as a request parameter. + * + * @param request current request + * @return true when no Roller session is present or the submitted salt is valid + */ + public static boolean consumeSubmittedSalt(HttpServletRequest request) { + RollerSession rollerSession = RollerSession.getRollerSession(request); + if (rollerSession == null) { + return true; + } + + String userId = rollerSession.getAuthenticatedUser() != null + ? rollerSession.getAuthenticatedUser().getId() : ""; + String salt = request.getParameter("salt"); + if (salt == null) { + return false; + } + + SaltCache saltCache = SaltCache.getInstance(); + synchronized (saltCache) { + if (!Objects.equals(saltCache.get(salt), userId)) { + return false; + } + saltCache.remove(salt); + } + return true; + } + + /** + * Returns true for a multipart form POST, which Struts parses after the + * servlet filters have run. + * + * @param request current request + * @return true for multipart/form-data POST requests + */ + public static boolean isMultipartFormPost(HttpServletRequest request) { + if (!"POST".equalsIgnoreCase(request.getMethod())) { + return false; + } + + String contentType = request.getContentType(); + if (contentType == null) { + return false; + } + + int parameterStart = contentType.indexOf(';'); + String mediaType = parameterStart >= 0 + ? contentType.substring(0, parameterStart) : contentType; + return MULTIPART_FORM_DATA.equals(mediaType.trim().toLowerCase(Locale.ENGLISH)); + } +} diff --git a/app/src/main/java/org/apache/roller/weblogger/ui/core/filters/ValidateSaltFilter.java b/app/src/main/java/org/apache/roller/weblogger/ui/core/filters/ValidateSaltFilter.java index 586bff185b..f1a90e156d 100644 --- a/app/src/main/java/org/apache/roller/weblogger/ui/core/filters/ValidateSaltFilter.java +++ b/app/src/main/java/org/apache/roller/weblogger/ui/core/filters/ValidateSaltFilter.java @@ -19,9 +19,6 @@ package org.apache.roller.weblogger.ui.core.filters; import java.io.IOException; -import java.util.Collections; -import java.util.Objects; -import java.util.Set; import javax.servlet.Filter; import javax.servlet.FilterChain; @@ -31,12 +28,8 @@ import javax.servlet.ServletResponse; import javax.servlet.http.HttpServletRequest; -import org.apache.commons.lang3.StringUtils; import org.apache.commons.logging.Log; import org.apache.commons.logging.LogFactory; -import org.apache.roller.weblogger.config.WebloggerConfig; -import org.apache.roller.weblogger.ui.rendering.util.cache.SaltCache; -import org.apache.roller.weblogger.ui.core.RollerSession; /** * Filter checks all POST request for presence of valid salt value and rejects those without @@ -44,40 +37,26 @@ */ public class ValidateSaltFilter implements Filter { private static final Log log = LogFactory.getLog(ValidateSaltFilter.class); - private Set ignored = Collections.emptySet(); @Override public void doFilter(ServletRequest request, ServletResponse response, FilterChain chain) throws IOException, ServletException { HttpServletRequest httpReq = (HttpServletRequest) request; - String requestURL = httpReq.getRequestURL().toString(); - String queryString = httpReq.getQueryString(); - if (queryString != null) { - requestURL += "?" + queryString; - } - - if ("POST".equals(httpReq.getMethod()) && !isIgnoredURL(requestURL)) { - RollerSession rollerSession = RollerSession.getRollerSession(httpReq); - if (rollerSession != null) { - String userId = rollerSession.getAuthenticatedUser() != null ? rollerSession.getAuthenticatedUser().getId() : ""; - - Object saltObject = httpReq.getAttribute("salt"); // multi-form post case - String salt = saltObject != null ? saltObject.toString() : null; - salt = salt != null ? salt : httpReq.getParameter("salt"); - SaltCache saltCache = SaltCache.getInstance(); - if (salt == null || !Objects.equals(saltCache.get(salt), userId)) { - if (log.isDebugEnabled()) { - log.debug("Valid salt value not found on POST to URL : " + httpReq.getServletPath()); - } - throw new ServletException("Security Violation"); - } + if ("POST".equalsIgnoreCase(httpReq.getMethod())) { + if (SaltValidator.isMultipartFormPost(httpReq) && isStrutsAction(httpReq)) { + // Struts makes multipart parameters available after its upload + // interceptor; ValidateSaltInterceptor handles these requests. + chain.doFilter(request, response); + return; + } - // Remove salt from cache after successful validation - saltCache.remove(salt); + if (!SaltValidator.consumeSubmittedSalt(httpReq)) { if (log.isDebugEnabled()) { - log.debug("Salt used and invalidated: " + salt); + log.debug("Valid salt value not found on POST to URL : " + + httpReq.getServletPath()); } + throw new ServletException("Security Violation"); } } @@ -86,20 +65,14 @@ public void doFilter(ServletRequest request, ServletResponse response, @Override public void init(FilterConfig filterConfig) throws ServletException { - String urls = WebloggerConfig.getProperty("salt.ignored.urls"); - ignored = Set.of(StringUtils.stripAll(StringUtils.split(urls, ","))); } @Override public void destroy() { } - /** - * Checks if this is an ignored url defined in the salt.ignored.urls property - * @param theUrl the url - * @return true, if is ignored resource - */ - private boolean isIgnoredURL(String theUrl) { - return ignored.contains(theUrl); + private boolean isStrutsAction(HttpServletRequest request) { + String servletPath = request.getServletPath(); + return servletPath != null && servletPath.endsWith(".rol"); } -} \ No newline at end of file +} diff --git a/app/src/main/java/org/apache/roller/weblogger/ui/struts2/util/ValidateSaltInterceptor.java b/app/src/main/java/org/apache/roller/weblogger/ui/struts2/util/ValidateSaltInterceptor.java new file mode 100644 index 0000000000..b64f103d6a --- /dev/null +++ b/app/src/main/java/org/apache/roller/weblogger/ui/struts2/util/ValidateSaltInterceptor.java @@ -0,0 +1,58 @@ +/* + * 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. For additional information regarding + * copyright in this work, please see the NOTICE file in the top level + * directory of this distribution. + */ + +package org.apache.roller.weblogger.ui.struts2.util; + +import javax.servlet.ServletException; +import javax.servlet.http.HttpServletRequest; + +import org.apache.commons.logging.Log; +import org.apache.commons.logging.LogFactory; +import org.apache.roller.weblogger.ui.core.filters.SaltValidator; +import org.apache.struts2.StrutsStatics; + +import com.opensymphony.xwork2.ActionContext; +import com.opensymphony.xwork2.ActionInvocation; +import com.opensymphony.xwork2.interceptor.AbstractInterceptor; + +/** + * Validates salts after Struts has parsed a multipart form request. + */ +public class ValidateSaltInterceptor extends AbstractInterceptor implements StrutsStatics { + + private static final long serialVersionUID = 2446434402795510394L; + private static final Log log = LogFactory.getLog(ValidateSaltInterceptor.class); + + @Override + public String intercept(ActionInvocation invocation) throws Exception { + ActionContext context = invocation.getInvocationContext(); + HttpServletRequest request = (HttpServletRequest) context.get(HTTP_REQUEST); + + if (SaltValidator.isMultipartFormPost(request) + && !SaltValidator.consumeSubmittedSalt(request)) { + if (log.isDebugEnabled()) { + log.debug("Valid salt value not found on multipart POST to URL : " + + request.getServletPath()); + } + throw new ServletException("Security Violation"); + } + + return invocation.invoke(); + } +} diff --git a/app/src/main/resources/org/apache/roller/weblogger/config/roller.properties b/app/src/main/resources/org/apache/roller/weblogger/config/roller.properties index d73e7f9ca1..b71a48689c 100644 --- a/app/src/main/resources/org/apache/roller/weblogger/config/roller.properties +++ b/app/src/main/resources/org/apache/roller/weblogger/config/roller.properties @@ -388,9 +388,6 @@ schemeenforcement.https.urls=/roller_j_security_check,\ # Ignored extensions otherwise we get SSL mixed content issues schemeenforcement.https.ignored=css,gif,png,js -# Ignored urls for salt. These are for multipart/form-data submissions as we do not get any parameters -salt.ignored.urls=mediaFileAdd!save.rol,mediaFileEdit!save.rol,bookmarksImport!save.rol - #--------------------------------------------------------------------- # LDAP authentication properties -- valid only if LDAP authentication # authentication.method via authentication.method setting. diff --git a/app/src/main/resources/struts.xml b/app/src/main/resources/struts.xml index cc94ba6588..51763b1b50 100644 --- a/app/src/main/resources/struts.xml +++ b/app/src/main/resources/struts.xml @@ -37,6 +37,8 @@ class="org.apache.roller.weblogger.ui.struts2.util.UISecurityInterceptor" /> + + diff --git a/app/src/main/webapp/WEB-INF/web.xml b/app/src/main/webapp/WEB-INF/web.xml index 0418832da1..746d20d267 100644 --- a/app/src/main/webapp/WEB-INF/web.xml +++ b/app/src/main/webapp/WEB-INF/web.xml @@ -142,15 +142,15 @@ - LoadSaltFilter + ValidateSaltFilter /roller-ui/* - REQUEST - FORWARD - ValidateSaltFilter + LoadSaltFilter /roller-ui/* + REQUEST + FORWARD diff --git a/app/src/test/java/org/apache/roller/weblogger/ui/core/filters/SaltConfigurationTest.java b/app/src/test/java/org/apache/roller/weblogger/ui/core/filters/SaltConfigurationTest.java new file mode 100644 index 0000000000..1c3a9431b8 --- /dev/null +++ b/app/src/test/java/org/apache/roller/weblogger/ui/core/filters/SaltConfigurationTest.java @@ -0,0 +1,85 @@ +/* + * 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. For additional information regarding + * copyright in this work, please see the NOTICE file in the top level + * directory of this distribution. + */ + +package org.apache.roller.weblogger.ui.core.filters; + +import java.io.IOException; +import java.io.InputStream; +import java.nio.charset.StandardCharsets; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.regex.Matcher; +import java.util.regex.Pattern; + +import org.junit.jupiter.api.Test; + +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertTrue; + +public class SaltConfigurationTest { + + @Test + public void testSubmittedSaltIsValidatedBeforeResponseSaltIsLoaded() throws Exception { + String webXml = Files.readString(Path.of("src/main/webapp/WEB-INF/web.xml")); + + int validateMapping = filterMappingPosition(webXml, "ValidateSaltFilter"); + int loadMapping = filterMappingPosition(webXml, "LoadSaltFilter"); + + assertTrue(validateMapping >= 0, "ValidateSaltFilter mapping is missing"); + assertTrue(loadMapping >= 0, "LoadSaltFilter mapping is missing"); + assertTrue(validateMapping < loadMapping, + "ValidateSaltFilter must run before LoadSaltFilter"); + } + + @Test + public void testMultipartSaltValidationImmediatelyFollowsUploadInterceptor() throws Exception { + String strutsXml = readResource("/struts.xml"); + + Pattern adjacentInterceptors = Pattern.compile( + "\\s*" + + ""); + + assertTrue(adjacentInterceptors.matcher(strutsXml).find(), + "ValidateSaltInterceptor must immediately follow the upload interceptor"); + } + + @Test + public void testConfigurableSaltBypassIsRemoved() throws Exception { + String properties = readResource( + "/org/apache/roller/weblogger/config/roller.properties"); + + assertFalse(properties.contains("salt.ignored.urls")); + } + + private int filterMappingPosition(String webXml, String filterName) { + Pattern pattern = Pattern.compile("\\s*" + + Pattern.quote(filterName) + ""); + Matcher matcher = pattern.matcher(webXml); + return matcher.find() ? matcher.start() : -1; + } + + private String readResource(String path) throws IOException { + try (InputStream stream = SaltConfigurationTest.class.getResourceAsStream(path)) { + if (stream == null) { + throw new IOException("Test resource not found: " + path); + } + return new String(stream.readAllBytes(), StandardCharsets.UTF_8); + } + } +} diff --git a/app/src/test/java/org/apache/roller/weblogger/ui/core/filters/ValidateSaltFilterTest.java b/app/src/test/java/org/apache/roller/weblogger/ui/core/filters/ValidateSaltFilterTest.java index ab866d080a..fcc5226f1d 100644 --- a/app/src/test/java/org/apache/roller/weblogger/ui/core/filters/ValidateSaltFilterTest.java +++ b/app/src/test/java/org/apache/roller/weblogger/ui/core/filters/ValidateSaltFilterTest.java @@ -1,17 +1,16 @@ package org.apache.roller.weblogger.ui.core.filters; -import org.apache.roller.weblogger.config.WebloggerConfig; import org.apache.roller.weblogger.pojos.User; import org.apache.roller.weblogger.ui.core.RollerSession; import org.apache.roller.weblogger.ui.rendering.util.cache.SaltCache; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; +import org.mockito.InOrder; import org.mockito.Mock; import org.mockito.MockedStatic; import org.mockito.MockitoAnnotations; import javax.servlet.FilterChain; -import javax.servlet.FilterConfig; import javax.servlet.ServletException; import javax.servlet.http.HttpServletRequest; import javax.servlet.http.HttpServletResponse; @@ -47,8 +46,6 @@ public void setUp() { @Test public void testDoFilterWithGetMethod() throws Exception { when(request.getMethod()).thenReturn("GET"); - StringBuffer requestURL = new StringBuffer("https://example.com/app/ignoredurl"); - when(request.getRequestURL()).thenReturn(requestURL); filter.doFilter(request, response, chain); @@ -67,9 +64,6 @@ public void testDoFilterWithPostMethodAndValidSalt() throws Exception { when(request.getParameter("salt")).thenReturn("validSalt"); when(saltCache.get("validSalt")).thenReturn("userId"); when(rollerSession.getAuthenticatedUser()).thenReturn(new TestUser("userId")); - StringBuffer requestURL = new StringBuffer("https://example.com/app/ignoredurl"); - when(request.getRequestURL()).thenReturn(requestURL); - filter.doFilter(request, response, chain); verify(chain).doFilter(request, response); @@ -88,9 +82,6 @@ public void testDoFilterWithPostMethodAndInvalidSalt() throws Exception { when(request.getMethod()).thenReturn("POST"); when(request.getParameter("salt")).thenReturn("invalidSalt"); when(saltCache.get("invalidSalt")).thenReturn(null); - StringBuffer requestURL = new StringBuffer("https://example.com/app/ignoredurl"); - when(request.getRequestURL()).thenReturn(requestURL); - assertThrows(ServletException.class, () -> { filter.doFilter(request, response, chain); }); @@ -109,9 +100,6 @@ public void testDoFilterWithPostMethodAndMismatchedUserId() throws Exception { when(request.getParameter("salt")).thenReturn("validSalt"); when(saltCache.get("validSalt")).thenReturn("differentUserId"); when(rollerSession.getAuthenticatedUser()).thenReturn(new TestUser("userId")); - StringBuffer requestURL = new StringBuffer("https://example.com/app/ignoredurl"); - when(request.getRequestURL()).thenReturn(requestURL); - assertThrows(ServletException.class, () -> { filter.doFilter(request, response, chain); }); @@ -129,9 +117,6 @@ public void testDoFilterWithPostMethodAndNullRollerSession() throws Exception { when(request.getMethod()).thenReturn("POST"); when(request.getParameter("salt")).thenReturn("validSalt"); when(saltCache.get("validSalt")).thenReturn(""); - StringBuffer requestURL = new StringBuffer("https://example.com/app/ignoredurl"); - when(request.getRequestURL()).thenReturn(requestURL); - filter.doFilter(request, response, chain); verify(saltCache, never()).remove("validSalt"); @@ -139,32 +124,110 @@ public void testDoFilterWithPostMethodAndNullRollerSession() throws Exception { } @Test - public void testDoFilterWithIgnoredURL() throws Exception { + public void testPostWithoutParameterRejectsRequestAttributeSalt() throws Exception { try (MockedStatic mockedRollerSession = mockStatic(RollerSession.class); - MockedStatic mockedSaltCache = mockStatic(SaltCache.class); - MockedStatic mockedWebloggerConfig = mockStatic(WebloggerConfig.class)) { + MockedStatic mockedSaltCache = mockStatic(SaltCache.class)) { mockedRollerSession.when(() -> RollerSession.getRollerSession(request)).thenReturn(rollerSession); mockedSaltCache.when(SaltCache::getInstance).thenReturn(saltCache); - mockedWebloggerConfig.when(() -> WebloggerConfig.getProperty("salt.ignored.urls")) - .thenReturn("https://example.com/app/ignoredurl?param1=value1&m2=value2"); when(request.getMethod()).thenReturn("POST"); - StringBuffer requestURL = new StringBuffer("https://example.com/app/ignoredurl"); - when(request.getRequestURL()).thenReturn(requestURL); - when(request.getQueryString()).thenReturn("param1=value1&m2=value2"); - when(request.getParameter("salt")).thenReturn(null); // No salt provided + when(request.getAttribute("salt")).thenReturn("responseSalt"); + when(request.getParameter("salt")).thenReturn(null); + when(rollerSession.getAuthenticatedUser()).thenReturn(new TestUser("userId")); + when(saltCache.get("responseSalt")).thenReturn("userId"); - filter.init(mock(FilterConfig.class)); - filter.doFilter(request, response, chain); + assertThrows(ServletException.class, + () -> filter.doFilter(request, response, chain)); - verify(chain).doFilter(request, response); + verify(chain, never()).doFilter(request, response); verify(saltCache, never()).get(anyString()); verify(saltCache, never()).remove(anyString()); } } + @Test + public void testSubmittedSaltCanOnlyBeUsedOnce() throws Exception { + try (MockedStatic mockedRollerSession = mockStatic(RollerSession.class); + MockedStatic mockedSaltCache = mockStatic(SaltCache.class)) { + + mockedRollerSession.when(() -> RollerSession.getRollerSession(request)).thenReturn(rollerSession); + mockedSaltCache.when(SaltCache::getInstance).thenReturn(saltCache); + + when(request.getMethod()).thenReturn("POST"); + when(request.getParameter("salt")).thenReturn("validSalt"); + when(rollerSession.getAuthenticatedUser()).thenReturn(new TestUser("userId")); + when(saltCache.get("validSalt")).thenReturn("userId", (String) null); + + filter.doFilter(request, response, chain); + assertThrows(ServletException.class, + () -> filter.doFilter(request, response, chain)); + + verify(chain, times(1)).doFilter(request, response); + verify(saltCache, times(1)).remove("validSalt"); + } + } + + @Test + public void testMultipartStrutsPostIsDeferred() throws Exception { + when(request.getMethod()).thenReturn("POST"); + when(request.getContentType()).thenReturn("multipart/form-data; boundary=abc123"); + when(request.getServletPath()).thenReturn("/roller-ui/mediaFileAdd!save.rol"); + + filter.doFilter(request, response, chain); + + verify(chain).doFilter(request, response); + verify(request, never()).getParameter("salt"); + } + + @Test + public void testMultipartNonStrutsPostIsNotDeferred() throws Exception { + try (MockedStatic mockedRollerSession = mockStatic(RollerSession.class)) { + mockedRollerSession.when(() -> RollerSession.getRollerSession(request)).thenReturn(rollerSession); + + when(request.getMethod()).thenReturn("POST"); + when(request.getContentType()).thenReturn("multipart/form-data; boundary=abc123"); + when(request.getServletPath()).thenReturn("/roller-ui/upload"); + when(request.getParameter("salt")).thenReturn(null); + + assertThrows(ServletException.class, + () -> filter.doFilter(request, response, chain)); + + verify(chain, never()).doFilter(request, response); + } + } + + @Test + public void testValidationRunsBeforeResponseSaltGeneration() throws Exception { + try (MockedStatic mockedRollerSession = mockStatic(RollerSession.class); + MockedStatic mockedSaltCache = mockStatic(SaltCache.class)) { + + mockedRollerSession.when(() -> RollerSession.getRollerSession(request)).thenReturn(rollerSession); + mockedSaltCache.when(SaltCache::getInstance).thenReturn(saltCache); + + when(request.getMethod()).thenReturn("POST"); + when(request.getParameter("salt")).thenReturn("submittedSalt"); + when(rollerSession.getAuthenticatedUser()).thenReturn(new TestUser("userId")); + when(saltCache.get("submittedSalt")).thenReturn("userId"); + + LoadSaltFilter loadSaltFilter = new LoadSaltFilter(); + FilterChain terminalChain = mock(FilterChain.class); + FilterChain loadSaltChain = (servletRequest, servletResponse) -> + loadSaltFilter.doFilter(servletRequest, servletResponse, terminalChain); + + filter.doFilter(request, response, loadSaltChain); + + InOrder order = inOrder(saltCache, request, terminalChain); + order.verify(saltCache).get("submittedSalt"); + order.verify(saltCache).remove("submittedSalt"); + order.verify(saltCache).put(anyString(), eq("userId")); + order.verify(request).setAttribute(eq("salt"), anyString()); + order.verify(terminalChain).doFilter(request, response); + } + } + private static class TestUser extends User { + private static final long serialVersionUID = 1L; private final String id; TestUser(String id) { diff --git a/app/src/test/java/org/apache/roller/weblogger/ui/struts2/util/ValidateSaltInterceptorTest.java b/app/src/test/java/org/apache/roller/weblogger/ui/struts2/util/ValidateSaltInterceptorTest.java new file mode 100644 index 0000000000..ba99bb2b01 --- /dev/null +++ b/app/src/test/java/org/apache/roller/weblogger/ui/struts2/util/ValidateSaltInterceptorTest.java @@ -0,0 +1,173 @@ +/* + * 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. For additional information regarding + * copyright in this work, please see the NOTICE file in the top level + * directory of this distribution. + */ + +package org.apache.roller.weblogger.ui.struts2.util; + +import javax.servlet.ServletException; +import javax.servlet.http.HttpServletRequest; + +import org.apache.roller.weblogger.pojos.User; +import org.apache.roller.weblogger.ui.core.RollerSession; +import org.apache.roller.weblogger.ui.rendering.util.cache.SaltCache; +import org.apache.struts2.StrutsStatics; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.mockito.Mock; +import org.mockito.MockedStatic; +import org.mockito.MockitoAnnotations; + +import com.opensymphony.xwork2.ActionContext; +import com.opensymphony.xwork2.ActionInvocation; + +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.mockito.Mockito.*; + +public class ValidateSaltInterceptorTest { + + private ValidateSaltInterceptor interceptor; + + @Mock + private ActionInvocation invocation; + + @Mock + private ActionContext context; + + @Mock + private HttpServletRequest request; + + @Mock + private RollerSession rollerSession; + + @Mock + private SaltCache saltCache; + + @BeforeEach + public void setUp() { + MockitoAnnotations.openMocks(this); + interceptor = new ValidateSaltInterceptor(); + when(invocation.getInvocationContext()).thenReturn(context); + when(context.get(StrutsStatics.HTTP_REQUEST)).thenReturn(request); + } + + @Test + public void testValidMultipartSaltIsConsumed() throws Exception { + try (MockedStatic mockedRollerSession = mockStatic(RollerSession.class); + MockedStatic mockedSaltCache = mockStatic(SaltCache.class)) { + mockedRollerSession.when(() -> RollerSession.getRollerSession(request)).thenReturn(rollerSession); + mockedSaltCache.when(SaltCache::getInstance).thenReturn(saltCache); + + configureMultipartPost("/roller-ui/mediaFileAdd!save.rol"); + when(request.getParameter("salt")).thenReturn("validSalt"); + when(rollerSession.getAuthenticatedUser()).thenReturn(new TestUser("userId")); + when(saltCache.get("validSalt")).thenReturn("userId"); + when(invocation.invoke()).thenReturn("success"); + + assertEquals("success", interceptor.intercept(invocation)); + + verify(saltCache).remove("validSalt"); + verify(invocation).invoke(); + } + } + + @Test + public void testMultipartPostWithoutSaltIsRejectedEvenWithResponseSaltAttribute() throws Exception { + try (MockedStatic mockedRollerSession = mockStatic(RollerSession.class)) { + mockedRollerSession.when(() -> RollerSession.getRollerSession(request)).thenReturn(rollerSession); + + configureMultipartPost("/roller-ui/bookmarksImport!save.rol"); + when(request.getParameter("salt")).thenReturn(null); + when(request.getAttribute("salt")).thenReturn("responseSalt"); + + assertThrows(ServletException.class, () -> interceptor.intercept(invocation)); + + verify(invocation, never()).invoke(); + } + } + + @Test + public void testInvalidMultipartSaltIsRejectedForAnyRollerAction() throws Exception { + try (MockedStatic mockedRollerSession = mockStatic(RollerSession.class); + MockedStatic mockedSaltCache = mockStatic(SaltCache.class)) { + mockedRollerSession.when(() -> RollerSession.getRollerSession(request)).thenReturn(rollerSession); + mockedSaltCache.when(SaltCache::getInstance).thenReturn(saltCache); + + configureMultipartPost("/roller-ui/arbitrary!save.rol"); + when(request.getParameter("salt")).thenReturn("invalidSalt"); + when(saltCache.get("invalidSalt")).thenReturn(null); + + assertThrows(ServletException.class, () -> interceptor.intercept(invocation)); + + verify(invocation, never()).invoke(); + verify(saltCache, never()).remove(anyString()); + } + } + + @Test + public void testMultipartSaltCannotBeReplayed() throws Exception { + try (MockedStatic mockedRollerSession = mockStatic(RollerSession.class); + MockedStatic mockedSaltCache = mockStatic(SaltCache.class)) { + mockedRollerSession.when(() -> RollerSession.getRollerSession(request)).thenReturn(rollerSession); + mockedSaltCache.when(SaltCache::getInstance).thenReturn(saltCache); + + configureMultipartPost("/roller-ui/mediaFileEdit!save.rol"); + when(request.getParameter("salt")).thenReturn("validSalt"); + when(rollerSession.getAuthenticatedUser()).thenReturn(new TestUser("userId")); + when(saltCache.get("validSalt")).thenReturn("userId", (String) null); + + interceptor.intercept(invocation); + assertThrows(ServletException.class, () -> interceptor.intercept(invocation)); + + verify(invocation, times(1)).invoke(); + verify(saltCache, times(1)).remove("validSalt"); + } + } + + @Test + public void testOrdinaryPostIsNotValidatedTwice() throws Exception { + when(request.getMethod()).thenReturn("POST"); + when(request.getContentType()).thenReturn("application/x-www-form-urlencoded"); + when(invocation.invoke()).thenReturn("success"); + + assertEquals("success", interceptor.intercept(invocation)); + + verify(request, never()).getParameter("salt"); + verify(invocation).invoke(); + } + + private void configureMultipartPost(String servletPath) { + when(request.getMethod()).thenReturn("POST"); + when(request.getContentType()).thenReturn("multipart/form-data; boundary=abc123"); + when(request.getServletPath()).thenReturn(servletPath); + } + + private static class TestUser extends User { + private static final long serialVersionUID = 1L; + private final String id; + + TestUser(String id) { + this.id = id; + } + + @Override + public String getId() { + return id; + } + } +}