diff --git a/lib/auth/discord-user-core.test.ts b/lib/auth/discord-user-core.test.ts new file mode 100644 index 0000000..a5b85ee --- /dev/null +++ b/lib/auth/discord-user-core.test.ts @@ -0,0 +1,131 @@ +import { describe, expect, test } from "bun:test"; +import { + reconcileDiscordUser, + type DiscordIdentity, + type DiscordUserOperations, +} from "@/lib/auth/discord-user-core"; + +type MemoryUser = { + id: string; + discordId: string; + name: string; + image: string | null; + email: string | null; +}; + +function createMemoryStore() { + const users = new Map(); + const userIdByDiscord = new Map(); + const accounts = new Map }>(); + let nextUserId = 1; + + const operations: DiscordUserOperations = { + async findLinkedUserId(discordId) { + return accounts.get(discordId)?.userId ?? null; + }, + async upsertUser(identity) { + let userId = userIdByDiscord.get(identity.discordId); + if (!userId) { + userId = `user-${nextUserId++}`; + userIdByDiscord.set(identity.discordId, userId); + users.set(userId, { + id: userId, + discordId: identity.discordId, + name: identity.displayName, + image: identity.avatarUrl, + email: identity.email ?? null, + }); + } + return userId; + }, + async insertAccount(userId, discordId, account) { + if (!accounts.has(discordId)) { + accounts.set(discordId, { + userId, + fields: { type: account?.type ?? "oauth" }, + }); + } + }, + async updateUser(userId, identity) { + const user = users.get(userId); + if (!user) return null; + const updated = { + ...user, + discordId: identity.discordId, + name: identity.displayName, + image: identity.avatarUrl, + ...(identity.email !== undefined ? { email: identity.email } : {}), + }; + users.set(userId, updated); + return updated; + }, + async updateAccount(discordId, update) { + const account = accounts.get(discordId); + if (account) account.fields = { ...account.fields, ...update }; + }, + }; + + return { users, accounts, operations }; +} + +const botIdentity: DiscordIdentity = { + discordId: "123456789012345678", + displayName: "Bot-created user", + avatarUrl: "https://cdn.discordapp.com/avatar.png", +}; + +describe("Discord account reconciliation", () => { + test("creates one user and one account for a new bot sender", async () => { + const store = createMemoryStore(); + const user = await reconcileDiscordUser(botIdentity, store.operations); + + expect(user.id).toBe("user-1"); + expect(store.users.size).toBe(1); + expect(store.accounts.size).toBe(1); + expect(store.accounts.get(botIdentity.discordId)?.userId).toBe(user.id); + }); + + test("concurrent first uploads resolve to the same database winner", async () => { + const store = createMemoryStore(); + const [first, second] = await Promise.all([ + reconcileDiscordUser(botIdentity, store.operations), + reconcileDiscordUser(botIdentity, store.operations), + ]); + + expect(first.id).toBe(second.id); + expect(store.users.size).toBe(1); + expect(store.accounts.size).toBe(1); + }); + + test("OAuth reuses a bot-created user and bot refreshes preserve email", async () => { + const store = createMemoryStore(); + const botUser = await reconcileDiscordUser(botIdentity, store.operations); + const oauthUser = await reconcileDiscordUser( + { + ...botIdentity, + displayName: "OAuth name", + email: "member@example.com", + account: { + type: "oauth", + access_token: "access-token", + refresh_token: "refresh-token", + }, + }, + store.operations, + ); + const refreshedByBot = await reconcileDiscordUser( + { ...botIdentity, displayName: "Latest Discord name" }, + store.operations, + ); + + expect(oauthUser.id).toBe(botUser.id); + expect(refreshedByBot.id).toBe(botUser.id); + expect(refreshedByBot.email).toBe("member@example.com"); + expect(store.users.size).toBe(1); + expect(store.accounts.size).toBe(1); + expect(store.accounts.get(botIdentity.discordId)?.fields).toMatchObject({ + access_token: "access-token", + refresh_token: "refresh-token", + }); + }); +}); diff --git a/lib/auth/discord-user-core.ts b/lib/auth/discord-user-core.ts new file mode 100644 index 0000000..af0a3cf --- /dev/null +++ b/lib/auth/discord-user-core.ts @@ -0,0 +1,78 @@ +import type { AdapterAccount } from "next-auth/adapters"; + +export type DiscordIdentity = { + discordId: string; + displayName: string; + avatarUrl: string | null; + email?: string | null; + account?: Partial; +}; + +export function discordAccountUpdate( + account: Partial | undefined, +) { + if (!account) return {}; + return { + ...(account.type !== undefined ? { type: account.type } : {}), + ...(account.refresh_token !== undefined + ? { refresh_token: account.refresh_token } + : {}), + ...(account.access_token !== undefined + ? { access_token: account.access_token } + : {}), + ...(account.expires_at !== undefined + ? { expires_at: account.expires_at } + : {}), + ...(account.token_type !== undefined + ? { token_type: account.token_type } + : {}), + ...(account.scope !== undefined ? { scope: account.scope } : {}), + ...(account.id_token !== undefined ? { id_token: account.id_token } : {}), + ...(account.session_state !== undefined + ? { session_state: account.session_state as string | null } + : {}), + }; +} + +export type DiscordAccountUpdate = ReturnType; + +export type DiscordUserOperations = { + findLinkedUserId(discordId: string): Promise; + upsertUser(identity: DiscordIdentity): Promise; + insertAccount( + userId: string, + discordId: string, + account: Partial | undefined, + ): Promise; + updateUser(userId: string, identity: DiscordIdentity): Promise; + updateAccount( + discordId: string, + update: DiscordAccountUpdate, + ): Promise; +}; + +export async function reconcileDiscordUser( + identity: DiscordIdentity, + operations: DiscordUserOperations, +) { + const { discordId, account } = identity; + let userId = await operations.findLinkedUserId(discordId); + + if (!userId) { + userId = await operations.upsertUser(identity); + await operations.insertAccount(userId, discordId, account); + + const winner = await operations.findLinkedUserId(discordId); + if (!winner) throw new Error("Failed to link Discord account"); + userId = winner; + } + + const user = await operations.updateUser(userId, identity); + const nextAccount = discordAccountUpdate(account); + if (Object.keys(nextAccount).length > 0) { + await operations.updateAccount(discordId, nextAccount); + } + + if (!user) throw new Error("Failed to reconcile Discord user"); + return user; +} diff --git a/lib/auth/discord-user.ts b/lib/auth/discord-user.ts index 7f99cd0..976dbeb 100644 --- a/lib/auth/discord-user.ts +++ b/lib/auth/discord-user.ts @@ -1,132 +1,94 @@ import "server-only"; -import type { AdapterAccount } from "next-auth/adapters"; import { and, eq } from "drizzle-orm"; import { db } from "@/db"; import { accounts, users } from "@/db/schema"; +import { + discordAccountUpdate, + reconcileDiscordUser, + type DiscordIdentity, +} from "@/lib/auth/discord-user-core"; -export type DiscordIdentity = { - discordId: string; - displayName: string; - avatarUrl: string | null; - email?: string | null; - account?: Partial; -}; - -function accountUpdate(account: Partial | undefined) { - if (!account) return {}; - return { - ...(account.type !== undefined ? { type: account.type } : {}), - ...(account.refresh_token !== undefined - ? { refresh_token: account.refresh_token } - : {}), - ...(account.access_token !== undefined - ? { access_token: account.access_token } - : {}), - ...(account.expires_at !== undefined - ? { expires_at: account.expires_at } - : {}), - ...(account.token_type !== undefined - ? { token_type: account.token_type } - : {}), - ...(account.scope !== undefined ? { scope: account.scope } : {}), - ...(account.id_token !== undefined ? { id_token: account.id_token } : {}), - ...(account.session_state !== undefined - ? { session_state: account.session_state as string | null } - : {}), - }; -} +export type { DiscordIdentity } from "@/lib/auth/discord-user-core"; export async function ensureDiscordUser( identity: DiscordIdentity ): Promise { - const { discordId, displayName, avatarUrl, email, account } = identity; - return db.transaction(async (tx) => { - const [linked] = await tx - .select({ userId: accounts.userId }) - .from(accounts) - .where( - and( - eq(accounts.provider, "discord"), - eq(accounts.providerAccountId, discordId) - ) - ) - .limit(1); - - let userId = linked?.userId; - if (!userId) { - const [user] = await tx - .insert(users) - .values({ - discordId, - name: displayName, - image: avatarUrl, - ...(email !== undefined ? { email } : {}), - }) - .onConflictDoUpdate({ - target: users.discordId, - set: { - name: displayName, - image: avatarUrl, - ...(email !== undefined ? { email } : {}), - }, - }) - .returning({ id: users.id }); - userId = user.id; - - await tx - .insert(accounts) - .values({ - userId, - provider: "discord", - providerAccountId: discordId, - type: account?.type ?? "oauth", - ...accountUpdate(account), - }) - .onConflictDoNothing({ - target: [accounts.provider, accounts.providerAccountId], - }); - - const [winner] = await tx + const findLinkedUserId = async (discordId: string) => { + const [linked] = await tx .select({ userId: accounts.userId }) .from(accounts) .where( and( eq(accounts.provider, "discord"), - eq(accounts.providerAccountId, discordId) - ) + eq(accounts.providerAccountId, discordId), + ), ) .limit(1); - if (!winner) throw new Error("Failed to link Discord account"); - userId = winner.userId; - } + return linked?.userId ?? null; + }; - const [user] = await tx - .update(users) - .set({ - discordId, - name: displayName, - image: avatarUrl, - ...(email !== undefined ? { email } : {}), - }) - .where(eq(users.id, userId)) - .returning(); - - const nextAccount = accountUpdate(account); - if (Object.keys(nextAccount).length > 0) { - await tx - .update(accounts) - .set(nextAccount) - .where( - and( - eq(accounts.provider, "discord"), - eq(accounts.providerAccountId, discordId) - ) - ); - } - - if (!user) throw new Error("Failed to reconcile Discord user"); - return user; + return reconcileDiscordUser(identity, { + findLinkedUserId, + async upsertUser({ discordId, displayName, avatarUrl, email }) { + const [user] = await tx + .insert(users) + .values({ + discordId, + name: displayName, + image: avatarUrl, + ...(email !== undefined ? { email } : {}), + }) + .onConflictDoUpdate({ + target: users.discordId, + set: { + name: displayName, + image: avatarUrl, + ...(email !== undefined ? { email } : {}), + }, + }) + .returning({ id: users.id }); + return user.id; + }, + async insertAccount(userId, discordId, account) { + await tx + .insert(accounts) + .values({ + userId, + provider: "discord", + providerAccountId: discordId, + type: account?.type ?? "oauth", + ...discordAccountUpdate(account), + }) + .onConflictDoNothing({ + target: [accounts.provider, accounts.providerAccountId], + }); + }, + async updateUser(userId, { discordId, displayName, avatarUrl, email }) { + const [user] = await tx + .update(users) + .set({ + discordId, + name: displayName, + image: avatarUrl, + ...(email !== undefined ? { email } : {}), + }) + .where(eq(users.id, userId)) + .returning(); + return user ?? null; + }, + async updateAccount(discordId, update) { + await tx + .update(accounts) + .set(update) + .where( + and( + eq(accounts.provider, "discord"), + eq(accounts.providerAccountId, discordId), + ), + ); + }, + }); }); }