Skip to content
Merged
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
6 changes: 3 additions & 3 deletions dist/main/index.js

Large diffs are not rendered by default.

150 changes: 123 additions & 27 deletions src/client/workload_identity_federation.ts
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,86 @@ 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 = 500;
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 responseMessage?: 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 candidate = err as { statusCode?: unknown; result?: unknown };
const status = candidate.statusCode;
if (typeof status === 'number') {
const result =
candidate.result && typeof candidate.result === 'object'
? (candidate.result as {
error_description?: unknown;
error?: { message?: unknown };
})
: undefined;
const responseMessage = result?.error_description || result?.error?.message;
return {
status,
errorClass: RETRYABLE_STS_STATUS_CODES.has(status)
? 'transient_http_response'
: 'non_retryable_http_response',
responseMessage:
typeof responseMessage === 'string'
? responseMessage.replace(/[\r\n]+/g, ' ').trim() || undefined
: undefined,
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 +138,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 +172,51 @@ 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, result: resp.result });
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) {
const responseMessage = failure.responseMessage
? `, response_message=${failure.responseMessage}`
: '';
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}` +
responseMessage,
);
}

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
209 changes: 207 additions & 2 deletions tests/client/workload_identity_client.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@
// See the License for the specific language governing permissions and
// limitations under the License.

import { test } from 'node:test';
import { test, TestContext } from 'node:test';
import assert from 'node:assert';

import { tmpdir } from 'os';
Expand All @@ -21,9 +21,214 @@ 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, result?: object): Error {
return Object.assign(new Error(`sensitive response body for ${statusCode}`), {
statusCode,
result,
});
}

function mockTimeouts(context: TestContext): number[] {
const delays: number[] = [];
context.mock.method(globalThis, 'setTimeout', ((callback: () => void, delay?: number) => {
delays.push(delay ?? 0);
callback();
return {} as NodeJS.Timeout;
}) as typeof setTimeout);
return delays;
}

test('#getToken retries transient STS responses', async (suite) => {
mockTimeouts(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 includes selected STS response details in the final error', async (suite) => {
await suite.test('uses error_description and normalizes line breaks', async () => {
const client = workloadIdentityClient();
mockTokenExchange(client, [
httpError(400, {
error_description: 'invalid subject token\nfor audience',
error: { message: 'nested message should not be used' },
}),
]);

await assert.rejects(client.getToken(), /response_message=invalid subject token for audience$/);
});

await suite.test('falls back to error.message', async () => {
const client = workloadIdentityClient();
mockTokenExchange(client, [httpError(400, { error: { message: 'nested STS error message' } })]);

await assert.rejects(client.getToken(), /response_message=nested STS error message$/);
});
});

test('#getToken emits sanitized attempt diagnostics', async (context) => {
mockTimeouts(context);
const logger = new RecordingLogger();
const client = workloadIdentityClient(logger);
const calls = mockTokenExchange(client, [
httpError(500, {
error_description: 'retryable STS error',
ignored: 'sensitive response body',
}),
httpError(400, {
error: { message: 'terminal STS error' },
ignored: 'sensitive response body',
}),
]);

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',
]);
assert.match(finalError, /response_message=terminal STS error$/);

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',
'retryable STS error',
]) {
assert.ok(!diagnostics.includes(secret), `diagnostics included ${secret}`);
}
});

test('#getToken bounds transient STS retries with exponential backoff', async (context) => {
const delays = mockTimeouts(context);
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);
assert.deepStrictEqual(delays, [500, 1000, 2000]);
});

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