Skip to content
Closed
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
2 changes: 1 addition & 1 deletion dist/main/index.js

Large diffs are not rendered by default.

132 changes: 105 additions & 27 deletions src/client/workload_identity_federation.ts
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,72 @@ import { errorMessage, writeSecureFile } from '@google-github-actions/actions-ut

import { AuthClient, Client, ClientParameters } from './client';

const STS_MAX_ATTEMPTS = 4;
const STS_RETRY_BACKOFF_MILLISECONDS = 100;
const RETRYABLE_STS_STATUS_CODES = new Set([408, 429, 500, 502, 503, 504]);
const RETRYABLE_CONNECTION_ERROR_CODES = new Set(['EAI_AGAIN', 'ECONNRESET', 'ETIMEDOUT']);

interface STSFailure {
readonly status?: number;
readonly errorClass: string;
readonly retryable: boolean;
}

function errorCode(err: unknown): string | undefined {
if (!err || typeof err !== 'object') {
return undefined;
}

const candidate = err as { code?: unknown; cause?: { code?: unknown } };
if (typeof candidate.code === 'string') {
return candidate.code;
}
if (typeof candidate.cause?.code === 'string') {
return candidate.cause.code;
}
return undefined;
}

function classifySTSFailure(err: unknown): STSFailure {
if (err && typeof err === 'object') {
const status = (err as { statusCode?: unknown }).statusCode;
if (typeof status === 'number') {
return {
status,
errorClass: RETRYABLE_STS_STATUS_CODES.has(status)
? 'transient_http_response'
: 'non_retryable_http_response',
retryable: RETRYABLE_STS_STATUS_CODES.has(status),
};
}
}

const code = errorCode(err);
if (code && RETRYABLE_CONNECTION_ERROR_CODES.has(code)) {
return {
errorClass: code,
retryable: true,
};
}

// @actions/http-client emits an uncoded error when its socket timeout fires.
if (err instanceof Error && err.message.startsWith('Request timeout:')) {
return {
errorClass: 'request_timeout',
retryable: true,
};
}

return {
errorClass: 'non_retryable_error',
retryable: false,
};
}

function sleep(milliseconds: number): Promise<void> {
return new Promise((resolve) => setTimeout(resolve, milliseconds));
}

/**
* WorkloadIdentityFederationClientParameters is used as input to the
* WorkloadIdentityFederationClient.
Expand Down Expand Up @@ -58,7 +124,6 @@ export class WorkloadIdentityFederationClient extends Client implements AuthClie

const iamHost = new URL(this._endpoints.iam).host;
this.#audience = `//${iamHost}/${this.#workloadIdentityProviderName}`;
this._logger.debug(`Computed audience`, this.#audience);
}

/**
Expand Down Expand Up @@ -93,34 +158,47 @@ export class WorkloadIdentityFederationClient extends Client implements AuthClie
subjectToken: this.#githubOIDCToken,
};

logger.debug(`Built request`, {
method: `POST`,
path: pth,
headers: headers,
body: body,
});

try {
const resp = await this._httpClient.postJson<{ access_token: string }>(pth, body, headers);
const statusCode = resp.statusCode || 500;
if (statusCode < 200 || statusCode > 299) {
throw new Error(`Failed to call ${pth}: HTTP ${statusCode}: ${resp.result || '[no body]'}`);
}

const result = resp.result;
if (!result) {
throw new Error(`Successfully called ${pth}, but the result was empty`);
const endpoint = new URL(pth).hostname;
for (let attempt = 1; attempt <= STS_MAX_ATTEMPTS; attempt++) {
try {
const resp = await this._httpClient.postJson<{ access_token: string }>(pth, body, headers);
const statusCode = resp.statusCode || 500;
if (statusCode < 200 || statusCode > 299) {
const err = new Error(`STS token exchange returned HTTP ${statusCode}`);
Object.assign(err, { statusCode });
throw err;
}

const result = resp.result;
if (!result) {
throw new Error(`STS token exchange returned an empty result`);
}

this.#cachedToken = result.access_token;
this.#cachedAt = now;
return result.access_token;
} catch (err) {
const failure = classifySTSFailure(err);
const status = failure.status ?? 'none';
logger.warning(
`STS request failed: operation=token_exchange, endpoint_class=${endpoint}, ` +
`status=${status}, error_class=${failure.errorClass}, ` +
`attempt=${attempt}/${STS_MAX_ATTEMPTS}`,
);

if (!failure.retryable || attempt === STS_MAX_ATTEMPTS) {
throw new Error(
`Failed to generate Google Cloud federated token: operation=token_exchange, ` +
`endpoint_class=${endpoint}, status=${status}, ` +
`error_class=${failure.errorClass}, attempt=${attempt}/${STS_MAX_ATTEMPTS}`,
);
}

await sleep(STS_RETRY_BACKOFF_MILLISECONDS * 2 ** (attempt - 1));
}

this.#cachedToken = result.access_token;
this.#cachedAt = now;
return result.access_token;
} catch (err) {
const msg = errorMessage(err);
throw new Error(
`Failed to generate Google Cloud federated token for ${this.#audience}: ${msg}`,
);
}

throw new Error(`STS token exchange failed unexpectedly`);
}

/**
Expand Down
158 changes: 157 additions & 1 deletion tests/client/workload_identity_client.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -21,9 +21,165 @@ import { readFileSync } from 'fs';

import { randomFilename } from '@google-github-actions/actions-utils';

import { NullLogger } from '../../src/logger';
import { Logger, NullLogger } from '../../src/logger';
import { WorkloadIdentityFederationClient } from '../../src/client/workload_identity_federation';

class RecordingLogger extends Logger {
readonly messages: string[] = [];

withNamespace(): Logger {
return this;
}

debug(...args: any[]) {
this.messages.push(args.join(' '));
}

warning(...args: any[]) {
this.messages.push(args.join(' '));
}
}

function workloadIdentityClient(
logger: Logger = new NullLogger(),
): WorkloadIdentityFederationClient {
return new WorkloadIdentityFederationClient({
logger,
universe: 'googleapis.com',
requestReason: 'sensitive-request-reason',
githubOIDCToken: 'sensitive-oidc-assertion',
githubOIDCTokenRequestURL: 'https://example.com/',
githubOIDCTokenRequestToken: 'sensitive-authorization-token',
githubOIDCTokenAudience: 'sensitive-audience',
workloadIdentityProviderName:
'projects/123/locations/global/workloadIdentityPools/pool/providers/provider',
serviceAccount: 'sensitive-service-account@example.com',
});
}

function mockTokenExchange(
client: WorkloadIdentityFederationClient,
outcomes: Array<object | Error>,
): () => number {
let calls = 0;
Object.defineProperty(client, '_httpClient', {
value: {
postJson: async () => {
const outcome = outcomes[calls++];
if (outcome instanceof Error) {
throw outcome;
}
return outcome;
},
},
});
return () => calls;
}

function httpError(statusCode: number): Error {
return Object.assign(new Error(`sensitive response body for ${statusCode}`), { statusCode });
}

test('#getToken retries transient STS responses', async (suite) => {
for (const statusCode of [408, 429, 500, 502, 503, 504]) {
await suite.test(`retries HTTP ${statusCode}`, async () => {
const client = workloadIdentityClient();
const calls = mockTokenExchange(client, [
httpError(statusCode),
{ statusCode: 200, result: { access_token: 'sensitive-access-token' } },
]);

assert.strictEqual(await client.getToken(), 'sensitive-access-token');
assert.strictEqual(calls(), 2);
});
}

for (const code of ['EAI_AGAIN', 'ECONNRESET', 'ETIMEDOUT']) {
await suite.test(`retries ${code}`, async () => {
const client = workloadIdentityClient();
const clientError = Object.assign(new Error('sensitive connection details'), { code });
const calls = mockTokenExchange(client, [
clientError,
{ statusCode: 200, result: { access_token: 'sensitive-access-token' } },
]);

assert.strictEqual(await client.getToken(), 'sensitive-access-token');
assert.strictEqual(calls(), 2);
});
}

await suite.test('retries the @actions/http-client socket timeout', async () => {
const client = workloadIdentityClient();
const calls = mockTokenExchange(client, [
new Error('Request timeout: /v1/token'),
{ statusCode: 200, result: { access_token: 'sensitive-access-token' } },
]);

assert.strictEqual(await client.getToken(), 'sensitive-access-token');
assert.strictEqual(calls(), 2);
});
});

test('#getToken does not retry permanent STS responses', async (suite) => {
for (const statusCode of [400, 401, 403]) {
await suite.test(`fails after HTTP ${statusCode}`, async () => {
const client = workloadIdentityClient();
const calls = mockTokenExchange(client, [httpError(statusCode)]);

await assert.rejects(client.getToken(), (err: Error) => {
assert.match(err.message, new RegExp(`status=${statusCode}`));
assert.match(err.message, /error_class=non_retryable_http_response/);
assert.match(err.message, /attempt=1\/4/);
return true;
});
assert.strictEqual(calls(), 1);
});
}
});

test('#getToken emits sanitized attempt diagnostics', async () => {
const logger = new RecordingLogger();
const client = workloadIdentityClient(logger);
const calls = mockTokenExchange(client, [httpError(500), httpError(400)]);

let finalError = '';
await assert.rejects(client.getToken(), (err: Error) => {
finalError = err.message;
return true;
});
assert.strictEqual(calls(), 2);
assert.deepStrictEqual(logger.messages, [
'STS request failed: operation=token_exchange, endpoint_class=sts.googleapis.com, status=500, error_class=transient_http_response, attempt=1/4',
'STS request failed: operation=token_exchange, endpoint_class=sts.googleapis.com, status=400, error_class=non_retryable_http_response, attempt=2/4',
]);

const diagnostics = [...logger.messages, finalError].join('\n');
for (const secret of [
'sensitive-oidc-assertion',
'sensitive-access-token',
'sensitive-authorization-token',
'sensitive-request-reason',
'sensitive-service-account@example.com',
'projects/123/locations/global/workloadIdentityPools/pool/providers/provider',
'sensitive response body',
]) {
assert.ok(!diagnostics.includes(secret), `diagnostics included ${secret}`);
}
});

test('#getToken bounds transient STS retries', async () => {
const client = workloadIdentityClient();
const calls = mockTokenExchange(client, [
httpError(503),
httpError(503),
httpError(503),
httpError(503),
]);

await assert.rejects(client.getToken(), /attempt=4\/4/);
assert.strictEqual(calls(), 4);
});

test('#createCredentialsFile', { concurrency: true }, async (suite) => {
await suite.test('writes the file', async () => {
const outputFile = pathjoin(tmpdir(), randomFilename());
Expand Down