fix(auth): bound all OIDC provider exchanges

This commit is contained in:
2026-08-17 10:03:07 +02:00
parent c86f01e886
commit 33ae7cfdc2
3 changed files with 437 additions and 59 deletions
+197 -56
View File
@@ -5,6 +5,7 @@ import {
customFetch,
discovery,
type Configuration,
type CustomFetch,
} from "openid-client";
import { constants, createPublicKey, verify as verifySignature } from "node:crypto";
import { parseConfiguredTransportUrl } from "./url-policy.js";
@@ -25,8 +26,17 @@ export interface OidcProtocol {
}
export class OidcProtocolError extends Error {
constructor(message = "oidc_protocol_invalid") {
super(message);
this.name = "OidcProtocolError";
}
}
/** A safe operational distinction for callers and diagnostics; provider details never cross this boundary. */
export class OidcProviderUnavailableError extends OidcProtocolError {
constructor() {
super("oidc_protocol_invalid");
super("oidc_provider_unavailable");
this.name = "OidcProviderUnavailableError";
}
}
@@ -38,13 +48,16 @@ export interface OidcProtocolOptions {
scopes: readonly string[];
groupsClaim: string;
fetch?: typeof globalThis.fetch;
httpTimeoutMs?: number;
jwksTimeoutMs?: number;
}
const MAX_GROUPS = 128;
const MAX_GROUP_LENGTH = 256;
const MAX_ID_TOKEN_LENGTH = 16 * 1024;
const MAX_JWKS_BYTES = 1024 * 1024;
const MAX_OIDC_RESPONSE_BYTES = 1024 * 1024;
const DEFAULT_HTTP_TIMEOUT_MS = 5_000;
const MAX_HTTP_TIMEOUT_MS = 30_000;
const DEFAULT_JWKS_TIMEOUT_MS = 5_000;
const MAX_JWKS_TIMEOUT_MS = 30_000;
const text = (value: unknown, maximum = 2048): value is string =>
@@ -166,70 +179,201 @@ function joseEcdsaSignatureToDer(signature: Buffer, partLength: number): Buffer
return Buffer.concat([Buffer.from([0x30]), derLength(sequence.length), sequence]);
}
async function readWithAbort(
reader: ReadableStreamDefaultReader<Uint8Array>,
signal: AbortSignal,
): Promise<ReadableStreamReadResult<Uint8Array>> {
signal.throwIfAborted();
return await new Promise((resolve, reject) => {
const aborted = () => {
cancelReaderBestEffort(reader);
reject(signal.reason);
};
signal.addEventListener("abort", aborted, { once: true });
reader.read().then(resolve, reject).finally(() => signal.removeEventListener("abort", aborted));
});
class OidcTransportError extends Error {
constructor(readonly availability: boolean) {
super(availability ? "oidc_provider_unavailable" : "oidc_provider_response_invalid");
this.name = "OidcTransportError";
}
}
interface BoundedOidcTransport {
customFetch: CustomFetch;
request(
input: RequestInfo | URL,
init?: RequestInit,
options?: { timeoutMs?: number; requireSuccess?: boolean },
): Promise<Response>;
}
function cancelReaderBestEffort(reader: ReadableStreamDefaultReader<Uint8Array>): void {
try {
void reader.cancel().catch(() => undefined);
void Promise.resolve(reader.cancel()).catch(() => undefined);
} catch {
// Cancellation is advisory; the bounded fetch timeout/rejection remains authoritative.
// Cancellation is advisory; the deadline and rejection remain authoritative.
}
}
async function boundedJwksBody(response: Response, signal: AbortSignal): Promise<string> {
if (!response.body) throw new OidcProtocolError();
const reader = response.body.getReader();
function cancelResponseBestEffort(response: Response): void {
if (!response.body) return;
try {
void Promise.resolve(response.body.cancel()).catch(() => undefined);
} catch {
// A late response is never allowed to turn an already-bounded request into an unhandled rejection.
}
}
function abortedReason(signal: AbortSignal): unknown {
return signal.reason ?? new DOMException("The operation was aborted", "AbortError");
}
async function awaitWithAbort<T>(
operation: Promise<T>,
signal: AbortSignal,
onLateResolution?: (value: T) => void,
): Promise<T> {
return await new Promise<T>((resolve, reject) => {
let settled = signal.aborted;
const cleanup = () => signal.removeEventListener("abort", aborted);
const aborted = () => {
if (settled) return;
settled = true;
cleanup();
reject(abortedReason(signal));
};
if (signal.aborted) reject(abortedReason(signal));
else signal.addEventListener("abort", aborted, { once: true });
operation.then(
(value) => {
if (settled) {
try { onLateResolution?.(value); } catch { /* best-effort late cleanup only */ }
return;
}
settled = true;
cleanup();
resolve(value);
},
(error: unknown) => {
if (settled) return;
settled = true;
cleanup();
reject(error);
},
);
});
}
function providerRequestUrl(input: RequestInfo | URL): URL {
const value = input instanceof URL ? input.href : input instanceof Request ? input.url : input;
return httpsEndpoint(value);
}
function declaredResponseLength(response: Response): number | undefined {
const declared = response.headers.get("content-length");
if (declared === null) return undefined;
if (!/^\d+$/.test(declared)) throw new OidcTransportError(false);
const length = Number(declared);
if (!Number.isSafeInteger(length) || length > MAX_OIDC_RESPONSE_BYTES) throw new OidcTransportError(false);
return length;
}
function safeBufferedResponse(response: Response, body: Buffer): Response {
const noBodyStatus = response.status === 204 || response.status === 205 || response.status === 304;
// Node accepts Buffer as a fetch body; the DOM declaration in this project does not model it.
return new Response(noBodyStatus ? null : body as unknown as BodyInit, {
status: response.status,
statusText: response.statusText,
headers: response.headers,
});
}
async function boundedResponse(
response: Response,
signal: AbortSignal,
requireSuccess: boolean,
): Promise<Response> {
const reader = response.body?.getReader();
const chunks: Buffer[] = [];
let total = 0;
let completed = false;
try {
if (!response.ok) throw new OidcProtocolError();
const declaredLength = response.headers.get("content-length");
if (declaredLength !== null) {
if (!/^\d+$/.test(declaredLength)) throw new OidcProtocolError();
const length = Number(declaredLength);
if (!Number.isSafeInteger(length) || length > MAX_JWKS_BYTES) throw new OidcProtocolError();
if (response.redirected || response.type === "opaqueredirect" || response.status >= 300 && response.status < 400) {
throw new OidcTransportError(false);
}
if (requireSuccess && !response.ok) throw new OidcTransportError(false);
declaredResponseLength(response);
if (!reader) {
completed = true;
return safeBufferedResponse(response, Buffer.alloc(0));
}
while (true) {
const { done, value } = await readWithAbort(reader, signal);
const { done, value } = await awaitWithAbort(reader.read(), signal);
if (done) break;
if (value.byteLength > MAX_JWKS_BYTES - total) throw new OidcProtocolError();
if (value.byteLength > MAX_OIDC_RESPONSE_BYTES - total) throw new OidcTransportError(false);
total += value.byteLength;
chunks.push(Buffer.from(value));
}
completed = true;
return safeBufferedResponse(response, Buffer.concat(chunks, total));
} finally {
try {
if (!completed) cancelReaderBestEffort(reader);
if (!completed && reader) cancelReaderBestEffort(reader);
} finally {
try { reader.releaseLock(); } catch { /* cancellation already made the response unusable */ }
try { reader?.releaseLock(); } catch { /* cancellation already made the response unusable */ }
}
}
signal.throwIfAborted();
try {
return new TextDecoder("utf-8", { fatal: true }).decode(Buffer.concat(chunks, total));
} catch {
throw new OidcProtocolError();
}
function createBoundedOidcTransport(
fetchImplementation: typeof globalThis.fetch,
defaultTimeoutMs: number,
): BoundedOidcTransport {
const request: BoundedOidcTransport["request"] = async (input, init, requestOptions) => {
let target: URL;
try {
target = providerRequestUrl(input);
} catch {
throw new OidcTransportError(false);
}
const timeoutMs = requestOptions?.timeoutMs ?? defaultTimeoutMs;
const controller = new AbortController();
const timeout = setTimeout(() => controller.abort(), timeoutMs);
timeout.unref();
const signal = init?.signal ? AbortSignal.any([init.signal, controller.signal]) : controller.signal;
try {
const requestInit: RequestInit = { ...init, redirect: "manual", signal };
const response = await awaitWithAbort(
Promise.resolve().then(() => fetchImplementation(target, requestInit)),
signal,
cancelResponseBestEffort,
);
return await boundedResponse(response, signal, requestOptions?.requireSuccess === true);
} catch (error) {
if (error instanceof OidcTransportError) throw error;
if (signal.aborted) throw new OidcTransportError(true);
throw new OidcTransportError(true);
} finally {
clearTimeout(timeout);
}
};
return {
request,
// openid-client's FetchBody is Fetch-compatible, but comes from a distinct declaration graph.
customFetch: (url, options) => request(url, options as unknown as RequestInit),
};
}
function availabilityFailure(error: unknown): boolean {
let current = error;
const seen = new Set<object>();
for (let depth = 0; depth < 8; depth += 1) {
if (current instanceof OidcTransportError) return current.availability;
if (!current || typeof current !== "object" || seen.has(current)) return false;
seen.add(current);
current = (current as { cause?: unknown }).cause;
}
return false;
}
function protocolFailure(error: unknown): OidcProtocolError {
if (error instanceof OidcProtocolError) return error;
return availabilityFailure(error) ? new OidcProviderUnavailableError() : new OidcProtocolError();
}
async function verifyIdTokenSignature(
idToken: unknown,
config: Configuration,
options: OidcProtocolOptions,
transport: BoundedOidcTransport,
jwksTimeoutMs: number,
): Promise<void> {
if (!text(idToken, MAX_ID_TOKEN_LENGTH)) throw new OidcProtocolError();
const [protectedPart, payloadPart, signaturePart, extra] = idToken.split(".");
@@ -241,15 +385,13 @@ async function verifyIdTokenSignature(
|| !metadata.id_token_signing_alg_values_supported.includes(header.alg)
|| !text(metadata.jwks_uri, 2048)) throw new OidcProtocolError();
const jwksUrl = httpsEndpoint(metadata.jwks_uri);
const controller = new AbortController();
const timeout = setTimeout(() => controller.abort(), options.jwksTimeoutMs ?? DEFAULT_JWKS_TIMEOUT_MS);
timeout.unref();
let response: Response;
try {
response = await (options.fetch ?? globalThis.fetch)(jwksUrl, {
headers: { accept: "application/json" }, redirect: "error", signal: controller.signal,
});
const body = await boundedJwksBody(response, controller.signal);
const response = await transport.request(
jwksUrl,
{ headers: { accept: "application/json" }, redirect: "manual" },
{ timeoutMs: jwksTimeoutMs, requireSuccess: true },
);
const body = new TextDecoder("utf-8", { fatal: true }).decode(await response.arrayBuffer());
const parsed = JSON.parse(body) as { keys?: unknown };
if (!Array.isArray(parsed.keys) || parsed.keys.length === 0 || parsed.keys.length > 16) throw new OidcProtocolError();
const matching = parsed.keys.filter((key): key is Record<string, unknown> =>
@@ -273,10 +415,7 @@ async function verifyIdTokenSignature(
: verifySignature(algorithm.digest, Buffer.from(`${protectedPart}.${payloadPart}`), publicKey, normalizedSignature);
if (!verified) throw new OidcProtocolError();
} catch (error) {
if (error instanceof OidcProtocolError) throw error;
throw new OidcProtocolError();
} finally {
clearTimeout(timeout);
throw protocolFailure(error);
}
}
@@ -286,8 +425,13 @@ export function createOidcProtocol(options: OidcProtocolOptions): OidcProtocol {
if (!text(options.clientId, 512) || !text(options.clientSecret, 4096)
|| !text(options.groupsClaim, 128) || options.scopes.length === 0 || options.scopes.length > 16
|| options.scopes.some((scope) => !text(scope, 128))
|| (options.httpTimeoutMs !== undefined && (!Number.isSafeInteger(options.httpTimeoutMs)
|| options.httpTimeoutMs < 1 || options.httpTimeoutMs > MAX_HTTP_TIMEOUT_MS))
|| (options.jwksTimeoutMs !== undefined && (!Number.isSafeInteger(options.jwksTimeoutMs)
|| options.jwksTimeoutMs < 1 || options.jwksTimeoutMs > MAX_JWKS_TIMEOUT_MS))) throw new OidcProtocolError();
const httpTimeoutMs = options.httpTimeoutMs ?? DEFAULT_HTTP_TIMEOUT_MS;
const jwksTimeoutMs = options.jwksTimeoutMs ?? DEFAULT_JWKS_TIMEOUT_MS;
const transport = createBoundedOidcTransport(options.fetch ?? globalThis.fetch, httpTimeoutMs);
let discovered: Promise<Configuration> | undefined;
const configuration = async (): Promise<Configuration> => {
@@ -299,7 +443,7 @@ export function createOidcProtocol(options: OidcProtocolOptions): OidcProtocol {
options.clientId,
{ client_secret: options.clientSecret, redirect_uris: [callbackUrl.href], response_types: ["code"] },
undefined,
options.fetch ? { [customFetch]: options.fetch as never } : undefined,
{ [customFetch]: transport.customFetch, timeout: httpTimeoutMs / 1000 },
);
const metadata = config.serverMetadata();
if (metadata.issuer !== options.issuer) throw new OidcProtocolError();
@@ -308,8 +452,7 @@ export function createOidcProtocol(options: OidcProtocolOptions): OidcProtocol {
httpsEndpoint(metadata.jwks_uri);
return config;
} catch (error) {
if (error instanceof OidcProtocolError) throw error;
throw new OidcProtocolError();
throw protocolFailure(error);
}
})();
}
@@ -330,8 +473,7 @@ export function createOidcProtocol(options: OidcProtocolOptions): OidcProtocol {
code_challenge_method: "S256",
});
} catch (error) {
if (error instanceof OidcProtocolError) throw error;
throw new OidcProtocolError();
throw protocolFailure(error);
}
},
async callback(input) {
@@ -346,13 +488,12 @@ export function createOidcProtocol(options: OidcProtocolOptions): OidcProtocol {
pkceCodeVerifier: input.codeVerifier,
idTokenExpected: true,
});
await verifyIdTokenSignature(tokens.id_token, config, options);
await verifyIdTokenSignature(tokens.id_token, config, transport, jwksTimeoutMs);
const claims = tokens.claims();
if (!claims || Array.isArray(claims)) throw new OidcProtocolError();
return identityFromClaims(claims as Record<string, unknown>, options);
} catch (error) {
if (error instanceof OidcProtocolError) throw error;
throw new OidcProtocolError();
throw protocolFailure(error);
}
},
async diagnose(signal) {