diff --git a/org.springframework.integration.security/src/main/java/org/springframework/integration/security/SecurityContextAssociatingHandlerInterceptor.java b/org.springframework.integration.security/src/main/java/org/springframework/integration/security/SecurityContextAssociatingHandlerInterceptor.java index c4e17d4f62..5ebb38677f 100644 --- a/org.springframework.integration.security/src/main/java/org/springframework/integration/security/SecurityContextAssociatingHandlerInterceptor.java +++ b/org.springframework.integration.security/src/main/java/org/springframework/integration/security/SecurityContextAssociatingHandlerInterceptor.java @@ -30,6 +30,13 @@ import org.springframework.security.context.SecurityContextHolder; */ public class SecurityContextAssociatingHandlerInterceptor extends InterceptingMessageHandler { + /** + * One time only set the strategy to be stack based to allow use of direct channels where push and pop is required rather than set and clear + */ + static { + SecurityContextHolder.setStrategyName(StackBasedSecurityContextHolderStrategy.class.getName()); + } + public SecurityContextAssociatingHandlerInterceptor(MessageHandler target) { super(target); } diff --git a/org.springframework.integration.security/src/main/java/org/springframework/integration/security/StackBasedSecurityContextHolderStrategy.java b/org.springframework.integration.security/src/main/java/org/springframework/integration/security/StackBasedSecurityContextHolderStrategy.java new file mode 100644 index 0000000000..3cf3516afe --- /dev/null +++ b/org.springframework.integration.security/src/main/java/org/springframework/integration/security/StackBasedSecurityContextHolderStrategy.java @@ -0,0 +1,76 @@ +/* + * Copyright 2002-2008 the original author or authors. + * + * Licensed 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.springframework.integration.security; + +import java.util.LinkedList; + +import org.apache.commons.logging.Log; +import org.apache.commons.logging.LogFactory; +import org.springframework.security.context.SecurityContext; +import org.springframework.security.context.SecurityContextHolderStrategy; +import org.springframework.security.context.SecurityContextImpl; +import org.springframework.util.Assert; + +/** + * Covers scenarios where direct channels are used and ensures that an existing + * {@link SecurityContext} is not unintentionally cleared + * @author Jonas Partner + * + */ +public class StackBasedSecurityContextHolderStrategy implements SecurityContextHolderStrategy { + + private Log logger = LogFactory.getLog(getClass()); + + private static ThreadLocal> contextHolder = new ThreadLocal>(); + + public void clearContext() { + if (getStackForThread().size() > 0) { + SecurityContext ctx = getStackForThread().pop(); + logger.debug("Popped security context " + ctx); + } + } + + public SecurityContext getContext() { + if (getStackForThread().peek() == null) { + logger.debug("Pushed new blank security context"); + getStackForThread().push(new SecurityContextImpl()); + } + + return (SecurityContext) getStackForThread().peek(); + } + + public void setContext(SecurityContext context) { + Assert.notNull(context, "Only non-null SecurityContext instances are permitted"); + + getStackForThread().push(context); + logger.debug("Pushed context " + context); + } + + protected LinkedList getStackForThread() { + if (contextHolder.get() == null) { + contextHolder.set(new LinkedList()); + } + return contextHolder.get(); + } + + public void clearStack() { + if (contextHolder.get() != null) { + contextHolder.set(null); + } + } + +} diff --git a/org.springframework.integration.security/src/test/java/org/springframework/integration/security/SecurityContextAssociatingHandlerInterceptorTests.java b/org.springframework.integration.security/src/test/java/org/springframework/integration/security/SecurityContextAssociatingHandlerInterceptorTests.java index 1a84ee9a87..22049352fc 100644 --- a/org.springframework.integration.security/src/test/java/org/springframework/integration/security/SecurityContextAssociatingHandlerInterceptorTests.java +++ b/org.springframework.integration.security/src/test/java/org/springframework/integration/security/SecurityContextAssociatingHandlerInterceptorTests.java @@ -20,8 +20,10 @@ import static org.junit.Assert.assertEquals; import static org.junit.Assert.assertNotSame; import static org.junit.Assert.assertNull; +import org.junit.After; import org.junit.Test; +import org.springframework.integration.dispatcher.DirectChannel; import org.springframework.integration.handler.MessageHandler; import org.springframework.integration.message.Message; import org.springframework.integration.message.StringMessage; @@ -36,6 +38,11 @@ import org.springframework.security.context.SecurityContextHolder; */ public class SecurityContextAssociatingHandlerInterceptorTests { + @After + public void clearSecurityContext(){ + SecurityContextHolder.clearContext(); + } + @Test public void testMessageWithSecurityContext() { final StubSecurityContext securityContext = new StubSecurityContext(); @@ -97,6 +104,27 @@ public class SecurityContextAssociatingHandlerInterceptorTests { assertNull("Security context still present after handler returned", SecurityContextHolder.getContext().getAuthentication()); } + + @Test + public void testExistingSecurityContextIsNotCleared(){ + SecurityContextHolder.setStrategyName(StackBasedSecurityContextHolderStrategy.class.getName()); + final StubSecurityContext securityContext = new StubSecurityContext(); + SecurityContextHolder.setContext(securityContext); + + StringMessage message = new StringMessage("test"); + + final MessageHandler handler = new MessageHandler() { + public Message handle(Message message) { + SecurityContext associatedContext = SecurityContextHolder.getContext(); + assertEquals("Wrong security context", securityContext, associatedContext); + return null; + } + }; + SecurityContextAssociatingHandlerInterceptor associatingInterceptor = + new SecurityContextAssociatingHandlerInterceptor(handler); + associatingInterceptor.handle(message); + assertEquals("Security context no logner set", securityContext, SecurityContextHolder.getContext()); + } @SuppressWarnings("serial") diff --git a/org.springframework.integration.security/src/test/java/org/springframework/integration/security/SecurityContextPropagatingChannelInterceptorTests.java b/org.springframework.integration.security/src/test/java/org/springframework/integration/security/SecurityContextPropagatingChannelInterceptorTests.java index 4f3c9a8713..2f452244e2 100644 --- a/org.springframework.integration.security/src/test/java/org/springframework/integration/security/SecurityContextPropagatingChannelInterceptorTests.java +++ b/org.springframework.integration.security/src/test/java/org/springframework/integration/security/SecurityContextPropagatingChannelInterceptorTests.java @@ -23,7 +23,6 @@ import static org.junit.Assert.assertTrue; import org.junit.After; import org.junit.Before; import org.junit.Test; - import org.springframework.integration.channel.QueueChannel; import org.springframework.integration.message.MessageHeader; import org.springframework.integration.message.StringMessage; @@ -78,11 +77,15 @@ public class SecurityContextPropagatingChannelInterceptorTests { channel.send(message); message = (StringMessage) channel.receive(0); MessageHeader header = message.getHeader(); + + Authentication auth = SecurityContextHolder.getContext().getAuthentication(); + assertFalse("Security context header found when no security context existed.", header.getAttributeNames().contains(SecurityContextPropagatingChannelInterceptor.SECURITY_CONTEXT_HEADER_ATTRIBUTE)); } + private void associateContextWithThread(){ SecurityContextHolder.setContext(securityContext); } diff --git a/org.springframework.integration.security/src/test/java/org/springframework/integration/security/SecurityEnforcingChannelInterceptorTests.java b/org.springframework.integration.security/src/test/java/org/springframework/integration/security/SecurityEnforcingChannelInterceptorTests.java index b26829e313..02cc4d8c58 100644 --- a/org.springframework.integration.security/src/test/java/org/springframework/integration/security/SecurityEnforcingChannelInterceptorTests.java +++ b/org.springframework.integration.security/src/test/java/org/springframework/integration/security/SecurityEnforcingChannelInterceptorTests.java @@ -18,6 +18,7 @@ package org.springframework.integration.security; import static org.junit.Assert.assertEquals; +import org.junit.After; import org.junit.Before; import org.junit.Test; @@ -29,6 +30,7 @@ import org.springframework.security.Authentication; import org.springframework.security.ConfigAttribute; import org.springframework.security.ConfigAttributeDefinition; import org.springframework.security.InsufficientAuthenticationException; +import org.springframework.security.context.SecurityContextHolder; /** * @author Jonas Partner @@ -44,7 +46,11 @@ public class SecurityEnforcingChannelInterceptorTests { public void setUp() { channel = new QueueChannel(); } - + + @After + public void clearSecurityContext(){ + SecurityContextHolder.clearContext(); + } @Test(expected = AccessDeniedException.class) public void testSendSecuredAndAccessDenied() { diff --git a/org.springframework.integration.security/src/test/java/org/springframework/integration/security/config/SecureChannelsParserTests.java b/org.springframework.integration.security/src/test/java/org/springframework/integration/security/config/SecureChannelsParserTests.java index 0e0d03d2ca..cdf89863cb 100644 --- a/org.springframework.integration.security/src/test/java/org/springframework/integration/security/config/SecureChannelsParserTests.java +++ b/org.springframework.integration.security/src/test/java/org/springframework/integration/security/config/SecureChannelsParserTests.java @@ -54,6 +54,7 @@ public class SecureChannelsParserTests { if (applicationContext != null) { applicationContext.close(); } + SecurityContextHolder.clearContext(); }