From e690808c70c1f4af391972aea252f7f89aa0ec28 Mon Sep 17 00:00:00 2001 From: YuqiGuo105 <131561736+YuqiGuo105@users.noreply.github.com> Date: Thu, 13 Aug 2026 11:28:04 -0600 Subject: [PATCH] Fix refresh request consumer invocation for web identity credentials --- ...oleWithWebIdentityCredentialsProvider.java | 6 ++- ...ithWebIdentityCredentialsProviderTest.java | 42 +++++++++++++++++++ 2 files changed, 46 insertions(+), 2 deletions(-) diff --git a/services/sts/src/main/java/software/amazon/awssdk/services/sts/auth/StsAssumeRoleWithWebIdentityCredentialsProvider.java b/services/sts/src/main/java/software/amazon/awssdk/services/sts/auth/StsAssumeRoleWithWebIdentityCredentialsProvider.java index aad3bfed0566..3e15a18a3f08 100644 --- a/services/sts/src/main/java/software/amazon/awssdk/services/sts/auth/StsAssumeRoleWithWebIdentityCredentialsProvider.java +++ b/services/sts/src/main/java/software/amazon/awssdk/services/sts/auth/StsAssumeRoleWithWebIdentityCredentialsProvider.java @@ -143,10 +143,12 @@ public Builder refreshRequest(Supplier assumeR * Similar to {@link #refreshRequest(AssumeRoleWithWebIdentityRequest)}, but takes a lambda to configure a new * {@link AssumeRoleWithWebIdentityRequest.Builder}. This removes the need to called * {@link AssumeRoleWithWebIdentityRequest#builder()} and {@link AssumeRoleWithWebIdentityRequest.Builder#build()}. + * The lambda is invoked each time the credentials are refreshed. */ public Builder refreshRequest(Consumer assumeRoleWithWebIdentityRequest) { - return refreshRequest(AssumeRoleWithWebIdentityRequest.builder().applyMutation(assumeRoleWithWebIdentityRequest) - .build()); + return refreshRequest(() -> + AssumeRoleWithWebIdentityRequest.builder().applyMutation(assumeRoleWithWebIdentityRequest) + .build()); } /** diff --git a/services/sts/src/test/java/software/amazon/awssdk/services/sts/auth/StsAssumeRoleWithWebIdentityCredentialsProviderTest.java b/services/sts/src/test/java/software/amazon/awssdk/services/sts/auth/StsAssumeRoleWithWebIdentityCredentialsProviderTest.java index 8f1e1c4808c3..bb06ea0380c4 100644 --- a/services/sts/src/test/java/software/amazon/awssdk/services/sts/auth/StsAssumeRoleWithWebIdentityCredentialsProviderTest.java +++ b/services/sts/src/test/java/software/amazon/awssdk/services/sts/auth/StsAssumeRoleWithWebIdentityCredentialsProviderTest.java @@ -15,6 +15,17 @@ package software.amazon.awssdk.services.sts.auth; +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.times; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import java.time.Duration; +import java.time.Instant; +import java.util.concurrent.atomic.AtomicInteger; +import org.junit.jupiter.api.Test; +import org.mockito.ArgumentCaptor; import software.amazon.awssdk.core.useragent.BusinessMetricFeatureId; import software.amazon.awssdk.services.sts.StsClient; import software.amazon.awssdk.services.sts.auth.StsAssumeRoleWithWebIdentityCredentialsProvider.Builder; @@ -56,4 +67,35 @@ protected AssumeRoleWithWebIdentityResponse callClient(StsClient client, AssumeR protected String providerName() { return BusinessMetricFeatureId.CREDENTIALS_STS_ASSUME_ROLE_WEB_ID.value(); } + + @Test + public void refreshRequestConsumerIsInvokedForEachCredentialRefresh() { + Credentials credentials = Credentials.builder() + .accessKeyId("a") + .secretAccessKey("b") + .sessionToken("c") + .expiration(Instant.now().minus(Duration.ofSeconds(5))) + .build(); + when(stsClient.assumeRoleWithWebIdentity(any(AssumeRoleWithWebIdentityRequest.class))) + .thenReturn(getResponse(credentials)); + + AtomicInteger tokenNumber = new AtomicInteger(); + try (StsAssumeRoleWithWebIdentityCredentialsProvider credentialsProvider = + StsAssumeRoleWithWebIdentityCredentialsProvider.builder() + .stsClient(stsClient) + .refreshRequest(request -> + request.webIdentityToken("token-" + + tokenNumber.incrementAndGet())) + .build()) { + credentialsProvider.resolveCredentials(); + credentialsProvider.resolveCredentials(); + } + + ArgumentCaptor requestCaptor = + ArgumentCaptor.forClass(AssumeRoleWithWebIdentityRequest.class); + verify(stsClient, times(2)).assumeRoleWithWebIdentity(requestCaptor.capture()); + assertThat(requestCaptor.getAllValues()) + .extracting(AssumeRoleWithWebIdentityRequest::webIdentityToken) + .containsExactly("token-1", "token-2"); + } }