diff --git a/pulsar-broker-auth-sasl/src/main/java/org/apache/pulsar/broker/authentication/AuthenticationProviderSasl.java b/pulsar-broker-auth-sasl/src/main/java/org/apache/pulsar/broker/authentication/AuthenticationProviderSasl.java index 6dc2ad0293306..5e539275af842 100644 --- a/pulsar-broker-auth-sasl/src/main/java/org/apache/pulsar/broker/authentication/AuthenticationProviderSasl.java +++ b/pulsar-broker-auth-sasl/src/main/java/org/apache/pulsar/broker/authentication/AuthenticationProviderSasl.java @@ -35,6 +35,7 @@ import static org.apache.pulsar.common.sasl.SaslConstants.SASL_STATE_SERVER_CHECK_TOKEN; import com.github.benmanes.caffeine.cache.Cache; import com.github.benmanes.caffeine.cache.Caffeine; +import com.google.common.annotations.VisibleForTesting; import jakarta.servlet.http.HttpServletRequest; import jakarta.servlet.http.HttpServletResponse; import java.io.IOException; @@ -334,6 +335,16 @@ public boolean authenticateHttpRequest(HttpServletRequest request, HttpServletRe } } + @VisibleForTesting + Cache getAuthStates() { + return authStates; + } + + @VisibleForTesting + void setAuthStates(Cache authStates) { + this.authStates = authStates; + } + private String sanitizeHeaderValue(String value) { if (value == null) { return null; diff --git a/pulsar-broker-auth-sasl/src/test/java/org/apache/pulsar/broker/authentication/SaslAuthenticateTest.java b/pulsar-broker-auth-sasl/src/test/java/org/apache/pulsar/broker/authentication/SaslAuthenticateTest.java index bb8595dacedd8..24d895dffca45 100644 --- a/pulsar-broker-auth-sasl/src/test/java/org/apache/pulsar/broker/authentication/SaslAuthenticateTest.java +++ b/pulsar-broker-auth-sasl/src/test/java/org/apache/pulsar/broker/authentication/SaslAuthenticateTest.java @@ -26,6 +26,7 @@ import static org.testng.Assert.assertTrue; import static org.testng.Assert.fail; import com.github.benmanes.caffeine.cache.Cache; +import com.github.benmanes.caffeine.cache.Caffeine; import com.google.common.collect.ImmutableSet; import jakarta.servlet.http.HttpServletRequest; import jakarta.servlet.http.HttpServletResponse; @@ -42,6 +43,9 @@ import java.util.Map; import java.util.Properties; import java.util.Set; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; import java.util.concurrent.TimeUnit; import javax.security.auth.login.Configuration; import lombok.Cleanup; @@ -356,12 +360,11 @@ public void testSaslOnlyAuthFirstStage() throws Exception { } @Test - @SuppressWarnings("unchecked") public void testMaxInflightContext() throws Exception { @Cleanup AuthenticationProviderSasl saslServer = new AuthenticationProviderSasl(); HttpServletRequest servletRequest = mock(HttpServletRequest.class); - doReturn("Init").when(servletRequest).getHeader("State"); + doReturn(SaslConstants.SASL_STATE_CLIENT_INIT).when(servletRequest).getHeader(SaslConstants.SASL_HEADER_STATE); conf.setInflightSaslContextExpiryMs(Integer.MAX_VALUE); conf.setMaxInflightSaslContext(1); saslServer.initialize(AuthenticationProvider.Context.builder().config(conf).build()); @@ -370,14 +373,65 @@ public void testMaxInflightContext() throws Exception { AuthenticationDataProvider dataProvider = authSasl.getAuthData("localhost"); AuthData initData1 = dataProvider.authenticate(AuthData.INIT_AUTH_DATA); doReturn(Base64.getEncoder().encodeToString(initData1.getBytes())).when( - servletRequest).getHeader("SASL-Token"); - doReturn(String.valueOf(i)).when(servletRequest).getHeader("SASL-Server-ID"); + servletRequest).getHeader(SaslConstants.SASL_AUTH_TOKEN); + doReturn(String.valueOf(i)).when(servletRequest).getHeader(SaslConstants.SASL_STATE_SERVER); saslServer.authenticateHttpRequest(servletRequest, mock(HttpServletResponse.class)); } - Field field = AuthenticationProviderSasl.class.getDeclaredField("authStates"); - field.setAccessible(true); - Cache cache = (Cache) field.get(saslServer); + Cache cache = saslServer.getAuthStates(); //only 1 context was left in the memory - assertEquals(cache.asMap().size(), 1); + // Caffeine may perform size-based eviction asynchronously, so force maintenance before asserting. + cache.cleanUp(); + assertEquals(cache.asMap().size(), conf.getMaxInflightSaslContext()); + } + + @Test + public void testMaxInflightContextWithDelayedCaffeineMaintenance() throws Exception { + @Cleanup + AuthenticationProviderSasl saslServer = new AuthenticationProviderSasl(); + HttpServletRequest servletRequest = mock(HttpServletRequest.class); + doReturn(SaslConstants.SASL_STATE_CLIENT_INIT).when(servletRequest).getHeader(SaslConstants.SASL_HEADER_STATE); + conf.setInflightSaslContextExpiryMs(Integer.MAX_VALUE); + conf.setMaxInflightSaslContext(1); + saslServer.initialize(AuthenticationProvider.Context.builder().config(conf).build()); + + CountDownLatch maintenanceStarted = new CountDownLatch(1); + CountDownLatch allowMaintenance = new CountDownLatch(1); + @Cleanup("shutdownNow") + ExecutorService maintenanceExecutor = Executors.newSingleThreadExecutor(); + Cache delayedMaintenanceCache = Caffeine.newBuilder() + .maximumSize(conf.getMaxInflightSaslContext()) + .expireAfterWrite(conf.getInflightSaslContextExpiryMs(), TimeUnit.MILLISECONDS) + .executor(command -> maintenanceExecutor.execute(() -> { + maintenanceStarted.countDown(); + try { + allowMaintenance.await(); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + } + command.run(); + })) + .build(); + saslServer.setAuthStates(delayedMaintenanceCache); + + try { + for (int i = 0; i < 10; i++) { + AuthenticationDataProvider dataProvider = authSasl.getAuthData("localhost"); + AuthData initData1 = dataProvider.authenticate(AuthData.INIT_AUTH_DATA); + doReturn(Base64.getEncoder().encodeToString(initData1.getBytes())).when( + servletRequest).getHeader(SaslConstants.SASL_AUTH_TOKEN); + doReturn(String.valueOf(i)).when(servletRequest).getHeader(SaslConstants.SASL_STATE_SERVER); + saslServer.authenticateHttpRequest(servletRequest, mock(HttpServletResponse.class)); + } + assertTrue(maintenanceStarted.await(5, TimeUnit.SECONDS)); + + Cache cache = saslServer.getAuthStates(); + assertTrue(cache.asMap().size() > conf.getMaxInflightSaslContext()); + + // Caffeine may perform size-based eviction asynchronously, so force maintenance before asserting. + cache.cleanUp(); + assertEquals(cache.asMap().size(), conf.getMaxInflightSaslContext()); + } finally { + allowMaintenance.countDown(); + } } }