From 586a6c6e5cf44446acb4ab62a3ced392a5928d9f Mon Sep 17 00:00:00 2001 From: Rujun Chen Date: Fri, 28 Oct 2022 13:53:01 +0800 Subject: [PATCH] Make UserPrincipalManager#getRoles more robust. --- .../aad/filter/UserPrincipalManager.java | 21 +++++++++++++----- .../aad/filter/UserPrincipalManagerTests.java | 22 ++++++++++++++----- 2 files changed, 31 insertions(+), 12 deletions(-) diff --git a/sdk/spring/spring-cloud-azure-autoconfigure/src/main/java/com/azure/spring/cloud/autoconfigure/aad/filter/UserPrincipalManager.java b/sdk/spring/spring-cloud-azure-autoconfigure/src/main/java/com/azure/spring/cloud/autoconfigure/aad/filter/UserPrincipalManager.java index b82ac96e5685..a3d2b1a9a4ed 100644 --- a/sdk/spring/spring-cloud-azure-autoconfigure/src/main/java/com/azure/spring/cloud/autoconfigure/aad/filter/UserPrincipalManager.java +++ b/sdk/spring/spring-cloud-azure-autoconfigure/src/main/java/com/azure/spring/cloud/autoconfigure/aad/filter/UserPrincipalManager.java @@ -30,12 +30,13 @@ import java.net.MalformedURLException; import java.net.URL; import java.text.ParseException; -import java.util.Collection; +import java.util.Collections; import java.util.HashSet; import java.util.Optional; import java.util.Set; import java.util.stream.Collectors; import java.util.stream.Stream; +import java.util.stream.StreamSupport; /** * A user principal manager to load user info from JWT. @@ -153,11 +154,19 @@ public UserPrincipal buildUserPrincipal(String aadIssuedBearerToken) throws Pars } Set getRoles(JWTClaimsSet set) { - return Optional.of(set) - .map(p -> p.getClaim(AadJwtClaimNames.ROLES)) - .map(Collection.class::cast) - .map(Collection::stream) - .orElseGet(Stream::empty) + if (set == null) { + return Collections.emptySet(); + } + Object rolesClaim = set.getClaim(AadJwtClaimNames.ROLES); + if (rolesClaim == null) { + return Collections.emptySet(); + } + if (rolesClaim instanceof Iterable) { + return StreamSupport.stream(((Iterable) rolesClaim).spliterator(), false) + .map(Object::toString) + .collect(Collectors.toSet()); + } + return Stream.of(rolesClaim) .map(Object::toString) .collect(Collectors.toSet()); } diff --git a/sdk/spring/spring-cloud-azure-autoconfigure/src/test/java/com/azure/spring/cloud/autoconfigure/aad/filter/UserPrincipalManagerTests.java b/sdk/spring/spring-cloud-azure-autoconfigure/src/test/java/com/azure/spring/cloud/autoconfigure/aad/filter/UserPrincipalManagerTests.java index bd4abb3c98d5..7f2feaa60c4d 100644 --- a/sdk/spring/spring-cloud-azure-autoconfigure/src/test/java/com/azure/spring/cloud/autoconfigure/aad/filter/UserPrincipalManagerTests.java +++ b/sdk/spring/spring-cloud-azure-autoconfigure/src/test/java/com/azure/spring/cloud/autoconfigure/aad/filter/UserPrincipalManagerTests.java @@ -20,7 +20,10 @@ import java.nio.file.Paths; import java.security.cert.CertificateFactory; import java.security.cert.X509Certificate; +import java.util.ArrayList; import java.util.Arrays; +import java.util.Collection; +import java.util.HashSet; import java.util.Set; import java.util.stream.Stream; @@ -77,14 +80,21 @@ void nullIssuer() { } @Test - void testRolesExtracted() { + void getRolesTest() { + rolesExtractedAsExpected(null, new ArrayList<>()); + rolesExtractedAsExpected("role1", Arrays.asList("role1")); + rolesExtractedAsExpected(Arrays.asList("role1", "role2"), Arrays.asList("role1", "role2")); + rolesExtractedAsExpected(new HashSet<>(Arrays.asList("role1", "role2")), Arrays.asList("role1", "role2")); + } + + private void rolesExtractedAsExpected(Object rolesClaimValue, Collection expected) { JWTClaimsSet set = new JWTClaimsSet.Builder() - .claim("roles", Arrays.asList("role1", "role2")) + .claim("roles", rolesClaimValue) .build(); - Set result = new UserPrincipalManager(null).getRoles(set); - assertEquals(2, result.size()); - assertTrue(result.contains("role1")); - assertTrue(result.contains("role2")); + Set actual = new UserPrincipalManager(null).getRoles(set); + assertEquals(expected.size(), actual.size()); + assertTrue(expected.containsAll(actual)); + assertTrue(actual.containsAll(expected)); } private String readJwtValidIssuerTxt() {