diff --git a/rt/rs/security/oauth-parent/oauth2/src/main/java/org/apache/cxf/rs/security/oauth2/services/AbstractAccessTokenValidator.java b/rt/rs/security/oauth-parent/oauth2/src/main/java/org/apache/cxf/rs/security/oauth2/services/AbstractAccessTokenValidator.java index 30cb6c4aa18..e131e09ede0 100644 --- a/rt/rs/security/oauth-parent/oauth2/src/main/java/org/apache/cxf/rs/security/oauth2/services/AbstractAccessTokenValidator.java +++ b/rt/rs/security/oauth-parent/oauth2/src/main/java/org/apache/cxf/rs/security/oauth2/services/AbstractAccessTokenValidator.java @@ -55,11 +55,22 @@ public abstract class AbstractAccessTokenValidator { private OAuthDataProvider dataProvider; private int maxValidationDataCacheSize; - private ConcurrentHashMap accessTokenValidations = + private long validationDataCacheLifetime = 60L; + private ConcurrentHashMap accessTokenValidations = new ConcurrentHashMap<>(); private JoseJwtConsumer jwtTokenConsumer; private boolean persistJwtEncoding = true; + private static final class AccessTokenValidationCacheEntry { + private final AccessTokenValidation validation; + private final long cachedAt; + + private AccessTokenValidationCacheEntry(AccessTokenValidation validation) { + this.validation = validation; + this.cachedAt = OAuthUtils.getIssuedAt(); + } + } + public void setTokenValidator(AccessTokenValidator validator) { setTokenValidators(Collections.singletonList(validator)); } @@ -105,8 +116,16 @@ protected AccessTokenValidation getAccessTokenValidation(String authScheme, Stri } AccessTokenValidation accessTokenV = null; + AccessTokenValidationCacheEntry cacheEntry = null; if (maxValidationDataCacheSize > 0) { - accessTokenV = accessTokenValidations.get(authSchemeData); + cacheEntry = accessTokenValidations.get(authSchemeData); + if (cacheEntry != null) { + if (isValidationDataCacheEntryExpired(cacheEntry)) { + accessTokenValidations.remove(authSchemeData, cacheEntry); + } else { + accessTokenV = cacheEntry.validation; + } + } } ServerAccessToken localAccessToken = null; if (accessTokenV == null) { @@ -150,6 +169,8 @@ protected AccessTokenValidation getAccessTokenValidation(String authScheme, Stri if (OAuthUtils.isExpired(accessTokenV.getTokenIssuedAt(), accessTokenV.getTokenLifetime())) { if (localAccessToken != null) { removeAccessToken(localAccessToken); + } else if (cacheEntry != null) { + accessTokenValidations.remove(authSchemeData, cacheEntry); } else if (maxValidationDataCacheSize > 0) { accessTokenValidations.remove(authSchemeData); } @@ -161,16 +182,21 @@ protected AccessTokenValidation getAccessTokenValidation(String authScheme, Stri && accessTokenV.getTokenNotBefore() > System.currentTimeMillis() / 1000L) { AuthorizationUtils.throwAuthorizationFailure(supportedSchemes, realm); } - if (maxValidationDataCacheSize > 0) { + if (maxValidationDataCacheSize > 0 && accessTokenV.isInitialValidationSuccessful()) { if (accessTokenValidations.size() >= maxValidationDataCacheSize) { // or delete the ones expiring sooner than others, etc accessTokenValidations.clear(); } - accessTokenValidations.put(authSchemeData, accessTokenV); + accessTokenValidations.putIfAbsent(authSchemeData, new AccessTokenValidationCacheEntry(accessTokenV)); } return accessTokenV; } + private boolean isValidationDataCacheEntryExpired(AccessTokenValidationCacheEntry cacheEntry) { + return validationDataCacheLifetime <= 0 + || OAuthUtils.isExpired(cacheEntry.cachedAt, validationDataCacheLifetime); + } + protected void removeAccessToken(ServerAccessToken at) { dataProvider.revokeToken(at.getClient(), at.getTokenKey(), @@ -185,6 +211,10 @@ public void setMaxValidationDataCacheSize(int maxValidationDataCacheSize) { this.maxValidationDataCacheSize = maxValidationDataCacheSize; } + public void setValidationDataCacheLifetime(long validationDataCacheLifetime) { + this.validationDataCacheLifetime = validationDataCacheLifetime; + } + public JoseJwtConsumer getJwtTokenConsumer() { return jwtTokenConsumer; } diff --git a/rt/rs/security/oauth-parent/oauth2/src/test/java/org/apache/cxf/rs/security/oauth2/services/AbstractAccessTokenValidatorTest.java b/rt/rs/security/oauth-parent/oauth2/src/test/java/org/apache/cxf/rs/security/oauth2/services/AbstractAccessTokenValidatorTest.java new file mode 100644 index 00000000000..821441b8a36 --- /dev/null +++ b/rt/rs/security/oauth-parent/oauth2/src/test/java/org/apache/cxf/rs/security/oauth2/services/AbstractAccessTokenValidatorTest.java @@ -0,0 +1,179 @@ +/** + * 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.cxf.rs.security.oauth2.services; + +import java.util.Collections; +import java.util.List; + +import jakarta.ws.rs.NotAuthorizedException; +import jakarta.ws.rs.core.MultivaluedMap; +import org.apache.cxf.jaxrs.ext.MessageContext; +import org.apache.cxf.rs.security.oauth2.common.AccessTokenValidation; +import org.apache.cxf.rs.security.oauth2.provider.AccessTokenValidator; +import org.apache.cxf.rs.security.oauth2.provider.OAuthServiceException; +import org.apache.cxf.rs.security.oauth2.utils.OAuthConstants; +import org.apache.cxf.rs.security.oauth2.utils.OAuthUtils; + +import org.junit.Test; + +import static org.junit.Assert.assertEquals; +import static org.junit.Assert.assertSame; +import static org.junit.Assert.assertThrows; + +public class AbstractAccessTokenValidatorTest { + + @Test + public void testValidationDataCacheReusesSuccessfulValidationWithinLifetime() { + TestAccessTokenValidator validator = new TestAccessTokenValidator(true); + TestValidator accessTokenValidator = createAccessTokenValidator(validator); + + AccessTokenValidation firstValidation = accessTokenValidator.getValidation("token"); + AccessTokenValidation secondValidation = accessTokenValidator.getValidation("token"); + + assertSame(firstValidation, secondValidation); + assertEquals(1, validator.getValidationCount()); + } + + @Test + public void testValidationDataCacheLifetimeBoundsSuccessfulValidationReuse() { + TestAccessTokenValidator validator = new TestAccessTokenValidator(true); + TestValidator accessTokenValidator = createAccessTokenValidator(validator); + accessTokenValidator.setValidationDataCacheLifetime(0L); + + accessTokenValidator.getValidation("token"); + accessTokenValidator.getValidation("token"); + + assertEquals(2, validator.getValidationCount()); + } + + @Test + public void testValidationDataCacheDoesNotReuseUnsuccessfulValidation() { + TestAccessTokenValidator validator = new TestAccessTokenValidator(false); + TestValidator accessTokenValidator = createAccessTokenValidator(validator); + + accessTokenValidator.getValidation("token"); + accessTokenValidator.getValidation("token"); + + assertEquals(2, validator.getValidationCount()); + } + + @Test + public void testValidationDataCacheReusesTokenExpiringWithinDefaultLifetime() { + TestAccessTokenValidator validator = new TestAccessTokenValidator(true); + validator.setTokenLifetime(30L); + TestValidator accessTokenValidator = createAccessTokenValidator(validator); + + AccessTokenValidation firstValidation = accessTokenValidator.getValidation("token"); + AccessTokenValidation secondValidation = accessTokenValidator.getValidation("token"); + + assertSame(firstValidation, secondValidation); + assertEquals(1, validator.getValidationCount()); + } + + @Test + public void testValidationDataCacheRejectsExpiredTokenBeforeCacheLifetime() { + TestAccessTokenValidator validator = new TestAccessTokenValidator(true); + TestValidator accessTokenValidator = createAccessTokenValidator(validator); + + AccessTokenValidation validation = accessTokenValidator.getValidation("token"); + validation.setTokenIssuedAt(OAuthUtils.getIssuedAt() - 2L); + validation.setTokenLifetime(1L); + + assertThrows(NotAuthorizedException.class, () -> accessTokenValidator.getValidation("token")); + accessTokenValidator.getValidation("token"); + + assertEquals(2, validator.getValidationCount()); + } + + @Test + public void testValidationDataCacheRevalidatesBeforeTokenExpires() { + TestAccessTokenValidator validator = new TestAccessTokenValidator(true); + TestValidator accessTokenValidator = createAccessTokenValidator(validator); + accessTokenValidator.setValidationDataCacheLifetime(0L); + + accessTokenValidator.getValidation("token"); + accessTokenValidator.getValidation("token"); + + assertEquals(2, validator.getValidationCount()); + } + + @Test + public void testValidationDataCacheLifetimeBoundsTokenWithNoLifetime() { + TestAccessTokenValidator validator = new TestAccessTokenValidator(true); + validator.setTokenLifetime(0L); + TestValidator accessTokenValidator = createAccessTokenValidator(validator); + accessTokenValidator.setValidationDataCacheLifetime(0L); + + accessTokenValidator.getValidation("token"); + accessTokenValidator.getValidation("token"); + + assertEquals(2, validator.getValidationCount()); + } + + private static TestValidator createAccessTokenValidator(TestAccessTokenValidator validator) { + TestValidator accessTokenValidator = new TestValidator(); + accessTokenValidator.setTokenValidator(validator); + accessTokenValidator.setMaxValidationDataCacheSize(10); + return accessTokenValidator; + } + + private static final class TestValidator extends AbstractAccessTokenValidator { + private AccessTokenValidation getValidation(String authSchemeData) { + return getAccessTokenValidation(OAuthConstants.BEARER_AUTHORIZATION_SCHEME, authSchemeData, null); + } + } + + private static final class TestAccessTokenValidator implements AccessTokenValidator { + private final boolean validationSuccessful; + private long tokenLifetime = 3600L; + private int validationCount; + + private TestAccessTokenValidator(boolean validationSuccessful) { + this.validationSuccessful = validationSuccessful; + } + + @Override + public List getSupportedAuthorizationSchemes() { + return Collections.singletonList(OAuthConstants.BEARER_AUTHORIZATION_SCHEME); + } + + @Override + public AccessTokenValidation validateAccessToken(MessageContext mc, + String authScheme, + String authSchemeData, + MultivaluedMap extraProps) + throws OAuthServiceException { + validationCount++; + AccessTokenValidation validation = new AccessTokenValidation(); + validation.setInitialValidationSuccessful(validationSuccessful); + validation.setTokenKey(authSchemeData); + validation.setTokenIssuedAt(OAuthUtils.getIssuedAt()); + validation.setTokenLifetime(tokenLifetime); + return validation; + } + + private void setTokenLifetime(long tokenLifetime) { + this.tokenLifetime = tokenLifetime; + } + + private int getValidationCount() { + return validationCount; + } + } +} \ No newline at end of file