Add cache for User entities in OIDC auth flow (#2822)

* Add cache for User entities in OIDC auth flow

* refactor: Address review feedback

- Refactor database call into a single, reusable method
- Increase the default cache size to 200
- Remove .recordStats() and using spy for testing
- Split unit tests into separate implementation test that use Mockito spies instead of checking internal cache stats
This commit is contained in:
Nilay Shah
2025-09-12 07:43:32 +00:00
committed by GitHub
parent 732c30b359
commit 06299ccb86
5 changed files with 170 additions and 3 deletions
@@ -16,13 +16,20 @@ package google.registry.request.auth;
import static com.google.common.net.HttpHeaders.AUTHORIZATION;
import static com.google.common.truth.Truth.assertThat;
import static google.registry.config.RegistryConfig.getUserAuthCachingDuration;
import static google.registry.config.RegistryConfig.getUserAuthMaxCachedEntries;
import static google.registry.request.auth.AuthModule.BEARER_PREFIX;
import static google.registry.request.auth.AuthModule.IAP_HEADER_NAME;
import static google.registry.testing.DatabaseHelper.createAdminUser;
import static google.registry.testing.DatabaseHelper.persistResource;
import static org.mockito.ArgumentMatchers.any;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.never;
import static org.mockito.Mockito.spy;
import static org.mockito.Mockito.verify;
import static org.mockito.Mockito.when;
import com.github.benmanes.caffeine.cache.LoadingCache;
import com.google.api.client.googleapis.auth.oauth2.GoogleIdToken.Payload;
import com.google.api.client.json.webtoken.JsonWebSignature;
import com.google.api.client.json.webtoken.JsonWebSignature.Header;
@@ -33,7 +40,9 @@ import dagger.Component;
import dagger.Module;
import dagger.Provides;
import google.registry.config.CredentialModule.ApplicationDefaultCredential;
import google.registry.config.RegistryConfig;
import google.registry.config.RegistryConfig.Config;
import google.registry.model.CacheUtils;
import google.registry.model.console.GlobalRole;
import google.registry.model.console.User;
import google.registry.model.console.UserRoles;
@@ -44,6 +53,7 @@ import google.registry.request.auth.OidcTokenAuthenticationMechanism.RegularOidc
import google.registry.util.GoogleCredentialsBundle;
import jakarta.inject.Singleton;
import jakarta.servlet.http.HttpServletRequest;
import java.util.Optional;
import org.junit.jupiter.api.AfterEach;
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
@@ -54,6 +64,9 @@ public class OidcTokenAuthenticationMechanismTest {
private static final String rawToken = "this-token";
private static final String email = "user@email.test";
private static final String unknownEmail = "bad-guy@evil.real";
private static final String gaiaId = "gaia-id";
private static final ImmutableSet<String> serviceAccounts =
ImmutableSet.of("service@email.test", "email@service.goog");
@@ -75,6 +88,12 @@ public class OidcTokenAuthenticationMechanismTest {
@BeforeEach
void beforeEach() throws Exception {
// 1. Create a brand new cache.
LoadingCache<String, Optional<User>> testCache =
CacheUtils.newCacheBuilder(getUserAuthCachingDuration())
.maximumSize(getUserAuthMaxCachedEntries())
.build(OidcTokenAuthenticationMechanism::loadUser);
OidcTokenAuthenticationMechanism.setCacheForTesting(testCache);
payload.setEmail(email);
payload.setSubject(gaiaId);
user = createAdminUser(email);
@@ -154,7 +173,7 @@ public class OidcTokenAuthenticationMechanismTest {
@Test
void testAuthenticate_unknownEmailAddress() throws Exception {
payload.setEmail("bad-guy@evil.real");
payload.setEmail(unknownEmail);
authResult = authenticationMechanism.authenticate(request);
assertThat(authResult).isEqualTo(AuthResult.NOT_AUTHENTICATED);
}
@@ -189,6 +208,62 @@ public class OidcTokenAuthenticationMechanismTest {
authenticationMechanism = component.regularOidcAuthenticationMechanism();
}
@Test
void testAuthenticate_ExistentUser_isCached() {
// Arrange: Create a spy of the actual cache object.
// A spy calls the real methods of the object while allowing us to verify interactions.
LoadingCache<String, Optional<User>> spiedCache =
spy(OidcTokenAuthenticationMechanism.userCache);
OidcTokenAuthenticationMechanism.setCacheForTesting(spiedCache);
// Act: Call the authenticate method.
authenticationMechanism.authenticate(request);
// Assert: Verify that the cache's "get" method was called exactly once.
// This confirms the cache is being used without checking its internal stats.
verify(spiedCache).get(email);
}
@Test
void testAuthenticate_nonExistentUser_isCached() {
// Arrange: Use an email that is not in the test database.
payload.setEmail(unknownEmail);
LoadingCache<String, Optional<User>> spiedCache =
spy(OidcTokenAuthenticationMechanism.userCache);
OidcTokenAuthenticationMechanism.setCacheForTesting(spiedCache);
// Act: Call the authenticate method.
authenticationMechanism.authenticate(request);
// Assert: Verify that the cache's "get" method was called for the unverified email.
// This confirms that we attempted to look up the unknown user in the cache.
verify(spiedCache).get(unknownEmail);
}
@Test
void testAuthenticate_whenCacheIsDisabled_cacheIsNotUsed() {
// Arrange: Explicitly disable the cache and create a spy.
RegistryConfig.overrideIsUserAuthCachingEnabledForTesting(false);
LoadingCache<String, Optional<User>> spiedCache =
spy(OidcTokenAuthenticationMechanism.userCache);
OidcTokenAuthenticationMechanism.setCacheForTesting(spiedCache);
// Act: Authenticate the user.
AuthResult authResult = authenticationMechanism.authenticate(request);
// Assert: The authentication should still succeed because the code falls back
// to the direct database call.
assertThat(authResult.isAuthenticated()).isTrue();
// Assert: Crucially, verify that the cache's "get" method was NEVER called.
// This proves the cache was correctly bypassed.
verify(spiedCache, never()).get(any(String.class));
// Teardown: Restore the default setting for other tests.
RegistryConfig.overrideIsUserAuthCachingEnabledForTesting(true);
}
@Singleton
@Component(modules = {AuthModule.class, TestModule.class})
interface TestComponent {
@@ -234,4 +309,12 @@ public class OidcTokenAuthenticationMechanismTest {
return GoogleCredentialsBundle.create(GoogleCredentials.newBuilder().build());
}
}
private void reinitializeCache() {
OidcTokenAuthenticationMechanism.userCache =
CacheUtils.newCacheBuilder(getUserAuthCachingDuration())
.maximumSize(getUserAuthMaxCachedEntries())
.recordStats()
.build(OidcTokenAuthenticationMechanism::loadUser);
}
}