import type { Env } from "./bindings.js"; import { prefixedKvKey } from "./kv_keys.js"; const AUTH_SESSION_CACHE_NAMESPACE = "auth_session_cache"; const BEARER_TOKEN_PATTERN = /^[A-Za-z0-9._~+/=-]{32,4096}$/; const SUBJECT_ID_PATTERN = /^[a-zA-Z0-9._:-]{3,128}$/; const DEVICE_ID_PATTERN = /^[a-zA-Z0-9._:-]{3,128}$/; const BETTER_AUTH_SESSION_QUERY = ` SELECT session.id, session.userId, session.expiresAt, device_context.device_id AS deviceId FROM better_auth_session AS session LEFT JOIN better_auth_session_device_context AS device_context ON device_context.session_id = session.id WHERE session.token = ? `; export interface AuthContext { userId: string; sessionId: string; tokenHash: string; expiresAt: string; deviceId?: string; } interface BetterAuthSessionRow extends Record { id: unknown; userId: unknown; expiresAt: unknown; deviceId?: unknown; } export type AuthErrorCode = | "authorization_missing" | "authorization_invalid" | "session_not_found" | "session_expired"; export class AuthError extends Error { constructor(readonly code: AuthErrorCode) { super(code); this.name = "AuthError"; } } export class AuthSessionCacheSchemaError extends Error { constructor(message: string) { super(message); this.name = "AuthSessionCacheSchemaError"; } } export function authSessionCacheKvKey(environment: string, tokenHash: string): string { if (!/^[a-f0-9]{64}$/.test(tokenHash)) { throw new AuthSessionCacheSchemaError("token_hash_invalid"); } return `${prefixedKvKey(environment, AUTH_SESSION_CACHE_NAMESPACE)}:${tokenHash}`; } export async function authenticatedRateLimitKey( environment: string, route: string, request: Request, ): Promise { let token: string | null; try { token = bearerToken(request); } catch (error) { if (error instanceof AuthError) { return `${environment}:${route}:authorization_invalid`; } throw error; } if (token === null) { return `${environment}:${route}:anonymous`; } return `${environment}:${route}:bearer:${await sha256Hex(token)}`; } export async function readAuthContext( request: Request, env: Env, now: Date = new Date(), ): Promise { const token = bearerToken(request); if (token === null) { throw new AuthError("authorization_missing"); } const tokenHash = await sha256Hex(token); const cacheKey = authSessionCacheKvKey(env.ELY_ENVIRONMENT, tokenHash); const document = await env.ELY_KV.get(cacheKey); if (document === null) { return readBetterAuthSessionContext(env, token, tokenHash, now); } const session = parseAuthSessionCacheDocument(document, tokenHash); if (Date.parse(session.expiresAt) <= now.getTime()) { throw new AuthError("session_expired"); } return session; } async function readBetterAuthSessionContext( env: Env, token: string, tokenHash: string, now: Date, ): Promise { const row = await env.ELY_DB.prepare(BETTER_AUTH_SESSION_QUERY) .bind(token) .first(); if (row === null) { throw new AuthError("session_not_found"); } const session: AuthContext = { userId: subjectId(stringField(row, "userId"), "userId"), sessionId: subjectId(stringField(row, "id"), "session_id"), tokenHash, expiresAt: timestampField(row, "expiresAt"), }; const deviceId = optionalStringField(row, "deviceId"); if (deviceId !== undefined) { session.deviceId = deviceIdValue(deviceId); } if (Date.parse(session.expiresAt) <= now.getTime()) { throw new AuthError("session_expired"); } return session; } export async function authTokenHash(token: string): Promise { return sha256Hex(token); } function bearerToken(request: Request): string | null { const header = request.headers.get("authorization"); if (header === null) { return null; } const [scheme, token, extra] = header.trim().split(/\s+/); if (scheme !== "Bearer" || token === undefined || extra !== undefined) { throw new AuthError("authorization_invalid"); } if (!BEARER_TOKEN_PATTERN.test(token)) { throw new AuthError("authorization_invalid"); } return token; } function parseAuthSessionCacheDocument(value: string, tokenHash: string): AuthContext { let parsed: unknown; try { parsed = JSON.parse(value); } catch { throw new AuthSessionCacheSchemaError("auth_session_cache_json_invalid"); } if (!isRecord(parsed)) { throw new AuthSessionCacheSchemaError("auth_session_cache_must_be_object"); } assertOnlyFields(parsed, ["version", "user_id", "session_id", "device_id", "expires_at"]); if (parsed.version !== 1) { throw new AuthSessionCacheSchemaError("auth_session_cache_version_invalid"); } const context: AuthContext = { userId: subjectId(stringField(parsed, "user_id"), "user_id"), sessionId: subjectId(stringField(parsed, "session_id"), "session_id"), tokenHash, expiresAt: isoTimestamp(stringField(parsed, "expires_at"), "expires_at"), }; const deviceId = optionalStringField(parsed, "device_id"); if (deviceId !== undefined) { context.deviceId = deviceIdValue(deviceId); } return context; } function subjectId(value: string, label: string): string { if (!SUBJECT_ID_PATTERN.test(value)) { throw new AuthSessionCacheSchemaError(`${label}_invalid`); } return value; } function deviceIdValue(value: string): string { if (!DEVICE_ID_PATTERN.test(value)) { throw new AuthSessionCacheSchemaError("device_id_invalid"); } return value; } function isoTimestamp(value: string, label: string): string { const timestamp = Date.parse(value); if (!Number.isFinite(timestamp)) { throw new AuthSessionCacheSchemaError(`${label}_invalid`); } return new Date(timestamp).toISOString(); } function timestampField(value: Record, field: string): string { const fieldValue = value[field]; if (typeof fieldValue === "string" && fieldValue.trim() !== "") { return isoTimestamp(fieldValue, field); } if (typeof fieldValue === "number" && Number.isFinite(fieldValue)) { return new Date(fieldValue).toISOString(); } if (fieldValue instanceof Date) { return fieldValue.toISOString(); } throw new AuthSessionCacheSchemaError(`${field}_required`); } function stringField(value: Record, field: string): string { const fieldValue = value[field]; if (typeof fieldValue !== "string" || fieldValue.trim() === "") { throw new AuthSessionCacheSchemaError(`${field}_required`); } return fieldValue.trim(); } function optionalStringField(value: Record, field: string): string | undefined { const fieldValue = value[field]; if (fieldValue === undefined || fieldValue === null) { return undefined; } if (typeof fieldValue !== "string" || fieldValue.trim() === "") { throw new AuthSessionCacheSchemaError(`${field}_invalid`); } return fieldValue.trim(); } function assertOnlyFields(value: Record, fields: string[]): void { const allowed = new Set(fields); for (const field of Object.keys(value)) { if (!allowed.has(field)) { throw new AuthSessionCacheSchemaError(`unexpected_field:${field}`); } } } function isRecord(value: unknown): value is Record { return typeof value === "object" && value !== null && !Array.isArray(value); } async function sha256Hex(value: string): Promise { const bytes = new TextEncoder().encode(value); const digest = await crypto.subtle.digest("SHA-256", bytes); return [...new Uint8Array(digest)].map((byte) => byte.toString(16).padStart(2, "0")).join(""); }