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 @@ -120,6 +120,12 @@ public class AwsCredentials extends ExternalAccountCredentials {

@Override
public AccessToken refreshAccessToken() throws IOException {
if (getServiceAccountImpersonationUrl() != null) {
if (this.impersonatedCredentials == null) {
this.impersonatedCredentials = this.buildImpersonatedCredentials();
}
return this.impersonatedCredentials.refreshAccessToken();
}
Comment thread
lsirac marked this conversation as resolved.
StsTokenExchangeRequest.Builder stsTokenExchangeRequest =
StsTokenExchangeRequest.newBuilder(retrieveSubjectToken(), getSubjectTokenType())
.setAudience(getAudience());
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -165,10 +165,18 @@ void refreshAccessToken_withServiceAccountImpersonation() throws IOException {
assertEquals(
transportFactory.transport.getServiceAccountAccessToken(), accessToken.getTokenValue());

// 3 AWS metadata requests (region, role, credentials) + 1 STS + 1 IAM generateAccessToken.
assertEquals(5, transportFactory.transport.getRequests().size());

// Validate metrics header is set correctly on the sts request.
Map<String, List<String>> headers =
transportFactory.transport.getRequests().get(6).getHeaders();
transportFactory.transport.getRequests().get(3).getHeaders();
ExternalAccountCredentialsTest.validateMetricsHeader(headers, "aws", true, false);

// Refreshing a second time reuses cached impersonatedCredentials and does not re-query AWS
// metadata while the source STS token is still unexpired.
awsCredential.refreshAccessToken();
assertEquals(6, transportFactory.transport.getRequests().size());
}

@Test
Expand Down Expand Up @@ -196,6 +204,7 @@ void refreshAccessToken_withServiceAccountImpersonationOptions() throws IOExcept

assertEquals(
transportFactory.transport.getServiceAccountAccessToken(), accessToken.getTokenValue());
assertEquals(5, transportFactory.transport.getRequests().size());

// Validate that default lifetime was set correctly on the request.
try (JsonParser jsonParser =
Expand All @@ -206,7 +215,7 @@ void refreshAccessToken_withServiceAccountImpersonationOptions() throws IOExcept

// Validate metrics header is set correctly on the sts request.
Map<String, List<String>> headers =
transportFactory.transport.getRequests().get(6).getHeaders();
transportFactory.transport.getRequests().get(3).getHeaders();
ExternalAccountCredentialsTest.validateMetricsHeader(headers, "aws", true, true);
}
}
Expand Down Expand Up @@ -246,7 +255,7 @@ void refreshAccessTokenProgrammaticRefresh_withServiceAccountImpersonation() thr

transportFactory.transport.setExpireTime(TestUtils.getDefaultExpireTime());

AwsSecurityCredentialsSupplier supplier =
TestAwsSecurityCredentialsSupplier supplier =
new TestAwsSecurityCredentialsSupplier("test", programmaticAwsCreds, null, null);

AwsCredentials awsCredential =
Expand All @@ -264,11 +273,17 @@ void refreshAccessTokenProgrammaticRefresh_withServiceAccountImpersonation() thr

assertEquals(
transportFactory.transport.getServiceAccountAccessToken(), accessToken.getTokenValue());
assertEquals(1, supplier.getRegionCount());
assertEquals(1, supplier.getCredentialsCount());

// Validate metrics header is set correctly on the sts request.
Map<String, List<String>> headers =
transportFactory.transport.getRequests().get(0).getHeaders();
ExternalAccountCredentialsTest.validateMetricsHeader(headers, "programmatic", true, false);

awsCredential.refreshAccessToken();
assertEquals(1, supplier.getRegionCount());
assertEquals(1, supplier.getCredentialsCount());
}

@Test
Expand Down Expand Up @@ -1353,6 +1368,8 @@ static class TestAwsSecurityCredentialsSupplier implements AwsSecurityCredential
private final AwsSecurityCredentials credentials;
private final IOException credentialException;
private final ExternalAccountSupplierContext expectedContext;
private int getRegionCount = 0;
private int getCredentialsCount = 0;

TestAwsSecurityCredentialsSupplier(
String region,
Expand All @@ -1365,8 +1382,17 @@ static class TestAwsSecurityCredentialsSupplier implements AwsSecurityCredential
this.expectedContext = expectedContext;
}

int getRegionCount() {
return getRegionCount;
}

int getCredentialsCount() {
return getCredentialsCount;
}

@Override
public String getRegion(ExternalAccountSupplierContext context) throws IOException {
getRegionCount++;
if (expectedContext != null) {
assertEquals(expectedContext.getAudience(), context.getAudience());
assertEquals(expectedContext.getSubjectTokenType(), context.getSubjectTokenType());
Expand All @@ -1377,6 +1403,7 @@ public String getRegion(ExternalAccountSupplierContext context) throws IOExcepti
@Override
public AwsSecurityCredentials getCredentials(ExternalAccountSupplierContext context)
throws IOException {
getCredentialsCount++;
if (credentialException != null) {
throw credentialException;
}
Expand Down
Loading