diff --git a/clients/src/main/java/org/apache/kafka/common/security/auth/SaslExtensions.java b/clients/src/main/java/org/apache/kafka/common/security/auth/SaslExtensions.java index c129f1ec400f7..ca4c4df6079ba 100644 --- a/clients/src/main/java/org/apache/kafka/common/security/auth/SaslExtensions.java +++ b/clients/src/main/java/org/apache/kafka/common/security/auth/SaslExtensions.java @@ -19,15 +19,34 @@ import java.util.Collections; import java.util.HashMap; import java.util.Map; +import java.util.Set; +import java.util.StringJoiner; +import javax.security.auth.Subject; /** - * A simple immutable value object class holding customizable SASL extensions + * A simple immutable value object class holding customizable SASL extensions. + * + *

+ * + * Note on object identity and equality: SaslExtensions intentionally + * overrides the standard {@link #equals(Object)} and {@link #hashCode()} methods calling their + * respective {@link Object#equals(Object)} and {@link Object#hashCode()} implementations. In so + * doing, it provides equality only via reference identity and will not base equality on + * the underlying values of its {@link #extensionsMap extentions map}. + * + *

+ * + * The reason for this approach to equality is based off of the manner in which + * credentials are stored in a {@link Subject}. SaslExtensions are added to and + * removed from a {@link Subject} via its {@link Subject#getPublicCredentials() public credentials}. + * The public credentials are stored in a {@link Set} in the {@link Subject}, so object equality + * therefore becomes a concern. With shallow, reference-based equality, distinct + * SaslExtensions instances with the same map values can be considered unique. This is + * critical to operations like token refresh. + * + * See KAFKA-14062 for more detail. */ public class SaslExtensions { - /** - * An "empty" instance indicating no SASL extensions - */ - public static final SaslExtensions NO_SASL_EXTENSIONS = new SaslExtensions(Collections.emptyMap()); private final Map extensionsMap; public SaslExtensions(Map extensionsMap) { @@ -41,21 +60,59 @@ public Map map() { return extensionsMap; } + /** + * Creates an "empty" instance indicating no SASL extensions. Do not cache the result of + * this method call for use by multiple {@link Subject}s as the references need to be + * unique. + * + *

+ * + * See the class-level documentation for details. + * @return Unique, but empty, SaslExtensions instance + */ + @SuppressWarnings("unchecked") + public static SaslExtensions empty() { + // It's ok to re-use the EMPTY_MAP instance as the object equality is on the outer + // SaslExtensions reference. + return new SaslExtensions(Collections.EMPTY_MAP); + } + + /** + * Implements equals using the reference comparison implementation from + * {@link Object#equals(Object)}. + * + *

+ * + * See the class-level documentation for details. + * + * @param o Other object to compare + * @return True if o == this + */ @Override - public boolean equals(Object o) { - if (this == o) return true; - if (o == null || getClass() != o.getClass()) return false; - return extensionsMap.equals(((SaslExtensions) o).extensionsMap); + public final boolean equals(Object o) { + return super.equals(o); } + /** + * Implements hashCode using the native implementation from + * {@link Object#hashCode()}. + * + *

+ * + * See the class-level documentation for details. + * + * @return Hash code of instance + */ @Override - public String toString() { - return extensionsMap.toString(); + public final int hashCode() { + return super.hashCode(); } @Override - public int hashCode() { - return extensionsMap.hashCode(); + public String toString() { + return new StringJoiner(", ", SaslExtensions.class.getSimpleName() + "[", "]") + .add("extensionsMap=" + extensionsMap) + .toString(); } } diff --git a/clients/src/main/java/org/apache/kafka/common/security/auth/SaslExtensionsCallback.java b/clients/src/main/java/org/apache/kafka/common/security/auth/SaslExtensionsCallback.java index c5bd449e0cc08..f2010afda67e4 100644 --- a/clients/src/main/java/org/apache/kafka/common/security/auth/SaslExtensionsCallback.java +++ b/clients/src/main/java/org/apache/kafka/common/security/auth/SaslExtensionsCallback.java @@ -26,13 +26,13 @@ * in the SASL exchange. */ public class SaslExtensionsCallback implements Callback { - private SaslExtensions extensions = SaslExtensions.NO_SASL_EXTENSIONS; + private SaslExtensions extensions = SaslExtensions.empty(); /** * Returns always non-null {@link SaslExtensions} consisting of the extension * names and values that are sent by the client to the server in the initial * client SASL authentication message. The default value is - * {@link SaslExtensions#NO_SASL_EXTENSIONS} so that if this callback is + * {@link SaslExtensions#empty()} so that if this callback is * unhandled the client will see a non-null value. */ public SaslExtensions extensions() { diff --git a/clients/src/main/java/org/apache/kafka/common/security/oauthbearer/internals/OAuthBearerClientInitialResponse.java b/clients/src/main/java/org/apache/kafka/common/security/oauthbearer/internals/OAuthBearerClientInitialResponse.java index a356f0da3ddb9..52623ff9fd4f2 100644 --- a/clients/src/main/java/org/apache/kafka/common/security/oauthbearer/internals/OAuthBearerClientInitialResponse.java +++ b/clients/src/main/java/org/apache/kafka/common/security/oauthbearer/internals/OAuthBearerClientInitialResponse.java @@ -108,7 +108,7 @@ public OAuthBearerClientInitialResponse(String tokenValue, String authorizationI this.tokenValue = Objects.requireNonNull(tokenValue, "token value must not be null"); this.authorizationId = authorizationId == null ? "" : authorizationId; validateExtensions(extensions); - this.saslExtensions = extensions != null ? extensions : SaslExtensions.NO_SASL_EXTENSIONS; + this.saslExtensions = extensions != null ? extensions : SaslExtensions.empty(); } /** diff --git a/clients/src/test/java/org/apache/kafka/common/security/SaslExtensionsTest.java b/clients/src/test/java/org/apache/kafka/common/security/SaslExtensionsTest.java index 9acb78cf3ef50..085baf70d2a2a 100644 --- a/clients/src/test/java/org/apache/kafka/common/security/SaslExtensionsTest.java +++ b/clients/src/test/java/org/apache/kafka/common/security/SaslExtensionsTest.java @@ -16,6 +16,7 @@ */ package org.apache.kafka.common.security; +import java.util.Collections; import org.apache.kafka.common.security.auth.SaslExtensions; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; @@ -23,6 +24,7 @@ import java.util.HashMap; import java.util.Map; +import static org.junit.jupiter.api.Assertions.assertNotEquals; import static org.junit.jupiter.api.Assertions.assertNull; import static org.junit.jupiter.api.Assertions.assertThrows; @@ -50,4 +52,30 @@ public void testCannotAddValueToMapReferenceAndGetFromExtensions() { this.map.put("hello", "42"); assertNull(extensions.map().get("hello")); } + + /** + * Tests that even when using the same underlying values in the map, two {@link SaslExtensions} + * are considered unique. + * + * @see SaslExtensions class-level documentation + */ + @Test + public void testExtensionsWithEqualValuesAreUnique() { + // If the maps are distinct objects but have the same underlying values, the SaslExtension + // objects should still be unique. + assertNotEquals(new SaslExtensions(Collections.singletonMap("key", "value")), + new SaslExtensions(Collections.singletonMap("key", "value")), + "SaslExtensions with unique maps should be unique"); + + // If the maps are the same object (with the same underlying values), the SaslExtension + // objects should still be unique. + assertNotEquals(new SaslExtensions(map), + new SaslExtensions(map), + "SaslExtensions with duplicate maps should be unique"); + + // If the maps are empty, the SaslExtension objects should still be unique. + assertNotEquals(SaslExtensions.empty(), + SaslExtensions.empty(), + "SaslExtensions returned from SaslExtensions.empty() should be unique"); + } } diff --git a/clients/src/test/java/org/apache/kafka/common/security/oauthbearer/OAuthBearerLoginModuleTest.java b/clients/src/test/java/org/apache/kafka/common/security/oauthbearer/OAuthBearerLoginModuleTest.java index ea03ec5bfa339..0dabeab1f433f 100644 --- a/clients/src/test/java/org/apache/kafka/common/security/oauthbearer/OAuthBearerLoginModuleTest.java +++ b/clients/src/test/java/org/apache/kafka/common/security/oauthbearer/OAuthBearerLoginModuleTest.java @@ -127,8 +127,8 @@ public void login1Commit1Login2Commit2Logout1Login3Commit3Logout2() throws Login // Create callback handler OAuthBearerToken[] tokens = new OAuthBearerToken[] {mock(OAuthBearerToken.class), mock(OAuthBearerToken.class), mock(OAuthBearerToken.class)}; - SaslExtensions[] extensions = new SaslExtensions[] {mock(SaslExtensions.class), - mock(SaslExtensions.class), mock(SaslExtensions.class)}; + SaslExtensions[] extensions = new SaslExtensions[] {saslExtensions(), + saslExtensions(), saslExtensions()}; TestCallbackHandler testTokenCallbackHandler = new TestCallbackHandler(tokens, extensions); // Create login modules @@ -208,7 +208,6 @@ public void login1Commit1Login2Commit2Logout1Login3Commit3Logout2() throws Login assertSame(extensions[2], publicCredentials.iterator().next()); verifyNoInteractions((Object[]) tokens); - verifyNoInteractions((Object[]) extensions); } @Test @@ -224,8 +223,8 @@ public void login1Commit1Logout1Login2Commit2Logout2() throws LoginException { // Create callback handler OAuthBearerToken[] tokens = new OAuthBearerToken[] {mock(OAuthBearerToken.class), mock(OAuthBearerToken.class)}; - SaslExtensions[] extensions = new SaslExtensions[] {mock(SaslExtensions.class), - mock(SaslExtensions.class)}; + SaslExtensions[] extensions = new SaslExtensions[] {saslExtensions(), + saslExtensions()}; TestCallbackHandler testTokenCallbackHandler = new TestCallbackHandler(tokens, extensions); // Create login modules @@ -270,7 +269,6 @@ public void login1Commit1Logout1Login2Commit2Logout2() throws LoginException { assertEquals(0, publicCredentials.size()); verifyNoInteractions((Object[]) tokens); - verifyNoInteractions((Object[]) extensions); } @Test @@ -285,8 +283,8 @@ public void loginAbortLoginCommitLogout() throws LoginException { // Create callback handler OAuthBearerToken[] tokens = new OAuthBearerToken[] {mock(OAuthBearerToken.class), mock(OAuthBearerToken.class)}; - SaslExtensions[] extensions = new SaslExtensions[] {mock(SaslExtensions.class), - mock(SaslExtensions.class)}; + SaslExtensions[] extensions = new SaslExtensions[] {saslExtensions(), + saslExtensions()}; TestCallbackHandler testTokenCallbackHandler = new TestCallbackHandler(tokens, extensions); // Create login module @@ -322,7 +320,6 @@ public void loginAbortLoginCommitLogout() throws LoginException { assertEquals(0, publicCredentials.size()); verifyNoInteractions((Object[]) tokens); - verifyNoInteractions((Object[]) extensions); } @Test @@ -338,8 +335,8 @@ public void login1Commit1Login2Abort2Login3Commit3Logout3() throws LoginExceptio // Create callback handler OAuthBearerToken[] tokens = new OAuthBearerToken[] {mock(OAuthBearerToken.class), mock(OAuthBearerToken.class), mock(OAuthBearerToken.class)}; - SaslExtensions[] extensions = new SaslExtensions[] {mock(SaslExtensions.class), - mock(SaslExtensions.class), mock(SaslExtensions.class)}; + SaslExtensions[] extensions = new SaslExtensions[] {saslExtensions(), saslExtensions(), + saslExtensions()}; TestCallbackHandler testTokenCallbackHandler = new TestCallbackHandler(tokens, extensions); // Create login modules @@ -406,7 +403,6 @@ public void login1Commit1Login2Abort2Login3Commit3Logout3() throws LoginExceptio assertSame(extensions[2], publicCredentials.iterator().next()); verifyNoInteractions((Object[]) tokens); - verifyNoInteractions((Object[]) extensions); } /** @@ -436,4 +432,21 @@ public void commitDoesNotThrowOnUnsupportedExtensionsCallback() throws LoginExce verifyNoInteractions((Object[]) tokens); } + + /** + * We don't want to use mocks for our tests as we need to make sure to test + * {@link SaslExtensions}' {@link SaslExtensions#equals(Object)} and + * {@link SaslExtensions#hashCode()} methods. + * + *

+ * + * We need to make distinct calls to this method (vs. caching the result and reusing it + * multiple times) because we need to ensure the {@link SaslExtensions} instances are unique. + * This properly mimics the behavior that is used during the token refresh logic. + * + * @return Unique, newly-created {@link SaslExtensions} instance + */ + private SaslExtensions saslExtensions() { + return SaslExtensions.empty(); + } }