diff --git a/apps/web/__tests__/unit/verification-token.test.ts b/apps/web/__tests__/unit/verification-token.test.ts new file mode 100644 index 0000000000..48d35f11de --- /dev/null +++ b/apps/web/__tests__/unit/verification-token.test.ts @@ -0,0 +1,203 @@ +import type { SQL } from "drizzle-orm"; +import { MySqlDialect } from "drizzle-orm/mysql-core"; +import type { MySql2Database } from "drizzle-orm/mysql2"; +import { describe, expect, it } from "vitest"; +import { DrizzleAdapter } from "../../../../packages/database/auth/drizzle-adapter"; + +interface VerificationTokenRow { + identifier: string; + token: string; + expires: Date; +} + +const dialect = new MySqlDialect(); + +function createMockDb(initialRows: VerificationTokenRow[]) { + let table = [...initialRows]; + let lastDeleteQuery: { sql: string; params: unknown[] } | null = null; + + const db = { + select: () => ({ + from: () => ({ + where: (pred: unknown) => { + const query = dialect.sqlToQuery(pred as SQL); + const identifierParam = String(query.params[0] ?? "").toLowerCase(); + return { + limit: async () => + table + .filter( + (row) => row.identifier.toLowerCase() === identifierParam, + ) + .slice(0, 1), + }; + }, + }), + }), + delete: () => ({ + where: (pred: unknown) => { + const query = dialect.sqlToQuery(pred as SQL); + lastDeleteQuery = query; + const [identifierParam, tokenParam] = query.params; + const initialCount = table.length; + table = table.filter( + (row) => + !(row.identifier === identifierParam && row.token === tokenParam), + ); + const affectedRows = initialCount - table.length; + return Promise.resolve([{ affectedRows }]); + }, + }), + transaction: async (cb: (tx: unknown) => Promise) => cb(db), + getTable: () => table, + getLastDeleteQuery: () => lastDeleteQuery, + }; + + return db; +} + +describe("useVerificationToken", () => { + it("burns the token on wrong guess and returns null", async () => { + const mockDb = createMockDb([ + { + identifier: "user@example.com", + token: "123456", + expires: new Date(Date.now() + 600000), + }, + ]); + + const adapter = DrizzleAdapter(mockDb as unknown as MySql2Database); + const result = await adapter.useVerificationToken?.({ + identifier: "USER@example.com", + token: "999999", + }); + + expect(result).toBeNull(); + const deleteQuery = mockDb.getLastDeleteQuery(); + expect(deleteQuery).not.toBeNull(); + expect(deleteQuery?.sql).toContain( + "`verification_tokens`.`identifier` = ?", + ); + expect(deleteQuery?.sql).toContain("`verification_tokens`.`token` = ?"); + expect(deleteQuery?.params).toEqual(["user@example.com", "123456"]); + expect(mockDb.getTable()).toHaveLength(0); + }); + + it("returns token and invalidates it on correct guess", async () => { + const mockDb = createMockDb([ + { + identifier: "user@example.com", + token: "123456", + expires: new Date(Date.now() + 600000), + }, + ]); + + const adapter = DrizzleAdapter(mockDb as unknown as MySql2Database); + const result = await adapter.useVerificationToken?.({ + identifier: "USER@example.com", + token: "123456", + }); + + expect(result).not.toBeNull(); + expect(result?.identifier).toBe("user@example.com"); + expect(result?.token).toBe("123456"); + const deleteQuery = mockDb.getLastDeleteQuery(); + expect(deleteQuery).not.toBeNull(); + expect(deleteQuery?.sql).toContain( + "`verification_tokens`.`identifier` = ?", + ); + expect(deleteQuery?.sql).toContain("`verification_tokens`.`token` = ?"); + expect(deleteQuery?.params).toEqual(["user@example.com", "123456"]); + expect(mockDb.getTable()).toHaveLength(0); + }); + + it("returns null if token does not exist", async () => { + const mockDb = createMockDb([]); + + const adapter = DrizzleAdapter(mockDb as unknown as MySql2Database); + const result = await adapter.useVerificationToken?.({ + identifier: "nonexistent@example.com", + token: "123456", + }); + + expect(result).toBeNull(); + expect(mockDb.getLastDeleteQuery()).toBeNull(); + }); + + it("prevents race condition by checking affectedRows on token consumption", async () => { + let table = [ + { + identifier: "user@example.com", + token: "123456", + expires: new Date(Date.now() + 600000), + }, + ]; + + let firstDeleteDone = false; + const mockDb = { + select: () => ({ + from: () => ({ + where: () => ({ + limit: async () => table.slice(0, 1), + }), + }), + }), + delete: () => ({ + where: (pred: unknown) => { + const query = dialect.sqlToQuery(pred as SQL); + expect(query.sql).toContain("`verification_tokens`.`identifier` = ?"); + expect(query.sql).toContain("`verification_tokens`.`token` = ?"); + if (!firstDeleteDone) { + firstDeleteDone = true; + table = []; + return Promise.resolve([{ affectedRows: 1 }]); + } + return Promise.resolve([{ affectedRows: 0 }]); + }, + }), + transaction: async (cb: (tx: unknown) => Promise) => cb(mockDb), + } as unknown as MySql2Database; + + const adapter = DrizzleAdapter(mockDb); + + const firstResult = await adapter.useVerificationToken?.({ + identifier: "USER@example.com", + token: "123456", + }); + + expect(firstResult).not.toBeNull(); + expect(firstResult?.token).toBe("123456"); + + const secondResult = await adapter.useVerificationToken?.({ + identifier: "USER@example.com", + token: "123456", + }); + + expect(secondResult).toBeNull(); + }); + + it("deletes only the selected token instance and preserves replacement tokens for the same user", async () => { + const mockDb = createMockDb([ + { + identifier: "user@example.com", + token: "123456", + expires: new Date(Date.now() + 600000), + }, + { + identifier: "user@example.com", + token: "replacement_token", + expires: new Date(Date.now() + 600000), + }, + ]); + + const adapter = DrizzleAdapter(mockDb as unknown as MySql2Database); + const result = await adapter.useVerificationToken?.({ + identifier: "USER@example.com", + token: "999999", + }); + + expect(result).toBeNull(); + const remaining = mockDb.getTable(); + expect(remaining.some((r) => r.token === "123456")).toBe(false); + expect(remaining.some((r) => r.token === "replacement_token")).toBe(true); + }); +}); diff --git a/packages/database/auth/drizzle-adapter.ts b/packages/database/auth/drizzle-adapter.ts index 16f7e10453..fae458b5e0 100644 --- a/packages/database/auth/drizzle-adapter.ts +++ b/packages/database/auth/drizzle-adapter.ts @@ -70,6 +70,21 @@ async function hasLinkedAccount(db: MySql2Database, userId: User.UserId) { return !!linkedAccount; } +function getAffectedRows(result: unknown): number { + if (Array.isArray(result)) { + return ( + (result[0] as { affectedRows?: number } | undefined)?.affectedRows ?? 0 + ); + } + return ( + (result as { affectedRows?: number; rowsAffected?: number } | undefined) + ?.affectedRows ?? + (result as { affectedRows?: number; rowsAffected?: number } | undefined) + ?.rowsAffected ?? + 0 + ); +} + export function DrizzleAdapter( db: MySql2Database, options?: { getSsoIdentity: () => ValidatedSsoIdentity | null }, @@ -510,31 +525,51 @@ export function DrizzleAdapter( return row; }, async useVerificationToken({ identifier, token }) { - const rows = await db - .select() - .from(verificationTokens) - .where(eq(verificationTokens.token, token)) - .limit(1); - const row = rows[0]; - if (!row) { - console.warn("[useVerificationToken] No token found"); - return null; - } const normalizedIdentifier = identifier?.toLowerCase() ?? ""; - const storedIdentifier = row.identifier?.toLowerCase() ?? ""; - if (normalizedIdentifier !== storedIdentifier) { - console.warn("[useVerificationToken] Identifier mismatch"); - return null; - } - await db - .delete(verificationTokens) - .where( - and( - eq(verificationTokens.token, token), - eq(verificationTokens.identifier, row.identifier), - ), + + const execute = async (tx: typeof db) => { + const rows = await tx + .select() + .from(verificationTokens) + .where(eq(verificationTokens.identifier, normalizedIdentifier)) + .limit(1); + const row = rows[0]; + if (!row) { + console.warn("[useVerificationToken] No token found"); + return null; + } + const storedIdentifier = row.identifier?.toLowerCase() ?? ""; + + const result = await tx + .delete(verificationTokens) + .where( + and( + eq(verificationTokens.identifier, row.identifier), + eq(verificationTokens.token, row.token), + ), + ); + + if (getAffectedRows(result) === 0) { + console.warn( + "[useVerificationToken] Token already consumed or invalid during deletion.", + ); + return null; + } + + if (row.token !== token) { + console.warn("[useVerificationToken] Token mismatch"); + return null; + } + + return { ...row, identifier: storedIdentifier }; + }; + + if (typeof db.transaction === "function") { + return await db.transaction(async (tx) => + execute(tx as unknown as typeof db), ); - return { ...row, identifier: storedIdentifier }; + } + return await execute(db); }, }; }