Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -143,10 +143,12 @@ public Builder refreshRequest(Supplier<AssumeRoleWithWebIdentityRequest> 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.Builder> assumeRoleWithWebIdentityRequest) {
return refreshRequest(AssumeRoleWithWebIdentityRequest.builder().applyMutation(assumeRoleWithWebIdentityRequest)
.build());
return refreshRequest(() ->
AssumeRoleWithWebIdentityRequest.builder().applyMutation(assumeRoleWithWebIdentityRequest)
.build());
}

/**
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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<AssumeRoleWithWebIdentityRequest> requestCaptor =
ArgumentCaptor.forClass(AssumeRoleWithWebIdentityRequest.class);
verify(stsClient, times(2)).assumeRoleWithWebIdentity(requestCaptor.capture());
assertThat(requestCaptor.getAllValues())
.extracting(AssumeRoleWithWebIdentityRequest::webIdentityToken)
.containsExactly("token-1", "token-2");
}
}