diff --git a/cloudflare/src/auth.ts b/cloudflare/src/auth.ts index be29e9e..bb7f452 100644 --- a/cloudflare/src/auth.ts +++ b/cloudflare/src/auth.ts @@ -5,6 +5,11 @@ 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 id, userId, expiresAt + FROM better_auth_session + WHERE token = ? +`; export interface AuthContext { userId: string; @@ -14,6 +19,12 @@ export interface AuthContext { deviceId?: string; } +interface BetterAuthSessionRow extends Record { + id: unknown; + userId: unknown; + expiresAt: unknown; +} + export type AuthErrorCode = | "authorization_missing" | "authorization_invalid" @@ -75,7 +86,7 @@ export async function readAuthContext( const cacheKey = authSessionCacheKvKey(env.ELY_ENVIRONMENT, tokenHash); const document = await env.ELY_KV.get(cacheKey); if (document === null) { - throw new AuthError("session_not_found"); + return readBetterAuthSessionContext(env, token, tokenHash, now); } const session = parseAuthSessionCacheDocument(document, tokenHash); @@ -85,6 +96,31 @@ export async function readAuthContext( 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"), + }; + if (Date.parse(session.expiresAt) <= now.getTime()) { + throw new AuthError("session_expired"); + } + return session; +} + export async function authTokenHash(token: string): Promise { return sha256Hex(token); } @@ -156,6 +192,20 @@ function isoTimestamp(value: string, label: string): string { 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() === "") { diff --git a/cloudflare/tests/api_controls.test.ts b/cloudflare/tests/api_controls.test.ts index 00354b4..9ca8a07 100644 --- a/cloudflare/tests/api_controls.test.ts +++ b/cloudflare/tests/api_controls.test.ts @@ -196,6 +196,55 @@ describe("api controls", () => { ]); }); + it("uses Better Auth D1 sessions when the session cache is cold", async () => { + const tokenHash = await authTokenHash(ACCESS_TOKEN); + const auditEvents: ElyAnalyticsDataPoint[] = []; + const kvReads: string[] = []; + const rateLimitKeys: string[] = []; + const d1Queries: string[] = []; + const d1Binds: unknown[][] = []; + const response = await withAuthenticatedApiControls( + new Request("https://elydora.test/api/devices", { + headers: { authorization: `Bearer ${ACCESS_TOKEN}` }, + }), + testEnv({ + auditEvents, + kvReads, + rateLimitKeys, + d1: testD1Database({ + firstRows: [ + { id: "session-01", userId: "user-01", expiresAt: "2099-01-01T00:00:00.000Z" }, + ], + queries: d1Queries, + binds: d1Binds, + }), + }), + "devices.list", + ["GET"], + (context) => + Promise.resolve( + jsonResponse({ user_id: context.userId, session_id: context.sessionId }, 200), + ), + ); + + assert.equal(response.status, 200); + assert.deepEqual(await response.json(), { user_id: "user-01", session_id: "session-01" }); + assert.deepEqual(rateLimitKeys, [`local:devices.list:bearer:${tokenHash}`]); + assert.deepEqual(kvReads, [authSessionCacheKvKey("local", tokenHash)]); + assert.equal(d1Queries.length, 1); + assert.match(d1Queries[0] ?? "", /FROM better_auth_session/); + assert.deepEqual(d1Binds, [[ACCESS_TOKEN]]); + assert.deepEqual(auditEvents[0]?.blobs?.slice(0, 7), [ + "devices.list", + "GET", + "/api/devices", + "handled", + "", + "", + "user-01", + ]); + }); + it("rejects expired authenticated sessions", async () => { const tokenHash = await authTokenHash(ACCESS_TOKEN); const response = await withAuthenticatedApiControls( @@ -222,6 +271,7 @@ describe("api controls", () => { interface TestEnvOptions { auditEvents?: ElyAnalyticsDataPoint[]; + d1?: Env["ELY_DB"]; kvEntries?: [string, string][]; kvReads?: string[]; rateLimitKeys?: string[]; @@ -246,7 +296,7 @@ function testEnv(options: TestEnvOptions = {}): Env { ELY_ENVIRONMENT: "local", ELY_AUTH_BASE_URL: "https://elydora.test", ELY_AUTH_SECRET: "test-auth-secret-for-api-controls", - ELY_DB: testD1Database(), + ELY_DB: options.d1 ?? testD1Database(), ELY_KV: { get(key: string): Promise { options.kvReads?.push(key); @@ -298,10 +348,18 @@ function testR2Bucket(): Env["ELY_STORAGE"] { }; } -function testD1Database(): Env["ELY_DB"] { +interface TestD1DatabaseOptions { + binds?: unknown[][]; + firstRows?: unknown[]; + queries?: string[]; +} + +function testD1Database(options: TestD1DatabaseOptions = {}): Env["ELY_DB"] { + let firstIndex = 0; return { - prepare() { - return testD1PreparedStatement(); + prepare(query: string) { + options.queries?.push(query); + return testD1PreparedStatement(options, () => firstIndex++); }, batch() { return Promise.resolve([]); @@ -312,13 +370,17 @@ function testD1Database(): Env["ELY_DB"] { }; } -function testD1PreparedStatement(): ReturnType { +function testD1PreparedStatement( + options: TestD1DatabaseOptions, + nextFirstIndex: () => number, +): ReturnType { return { - bind() { + bind(...values: unknown[]) { + options.binds?.push(values); return this; }, - first() { - return Promise.resolve(null); + first() { + return Promise.resolve((options.firstRows?.[nextFirstIndex()] as T | undefined) ?? null); }, all() { return Promise.resolve({ results: [] });